packages feed

proarrow-0.1.0.0: src/Proarrow/Category/Instance/FinSet.hs

{-# LANGUAGE AllowAmbiguousTypes #-}

{- HLINT ignore "Use elemIndex" -}

-- | The skeleton of the category of __finite sets__: objects are natural numbers (@'FS' n@) and a
-- morphism @'FS' n '~>' 'FS' m@ is a function stored as its table, a length-@n@ vector of indices
-- below @m@. Distributive and cartesian closed (exponentials via the 'Exp' type family), with all
-- structure computed concretely.
module Proarrow.Category.Instance.FinSet where

import Data.Containers.ListUtils (nubOrd)
import Data.Data (Proxy (..))
import Data.Fin (Fin (..), fin0, fin1, split, weakenLeft, weakenRight)
import Data.IntMap qualified as IM
import Data.List qualified as List
import Data.Maybe (fromJust, fromMaybe, isNothing)
import Data.Type.Nat (Mult, Nat (..), Nat0, Nat1, Nat2, Plus, SNat (..), SNatI, snat)
import Data.Vec.Lazy
  ( Vec (..)
  , chunks
  , concat
  , concatMap
  , reifyList
  , repeat
  , tabulate
  , toList
  , universe
  , zipWith
  , (!)
  , (++)
  )
import Prelude (($))
import Prelude qualified as P

import Proarrow.Category.Monoidal (Monoidal (..), MonoidalProfunctor (..), SymMonoidal (..))
import Proarrow.Category.Monoidal.Closed (Closed (..))
import Proarrow.Category.Monoidal.CopyDiscard (CopyDiscard)
import Proarrow.Category.Monoidal.Distributive (Distributive (..))
import Proarrow.Category.Topos (ElementaryTopos, HasEpiMonoFactorization (..), HasSubobjectClassifier (..))
import Proarrow.Colimit.BinaryCoproduct (HasBinaryCoproducts (..))
import Proarrow.Colimit.Coequalizer (HasCoequalizers (..))
import Proarrow.Colimit.Initial (HasInitialObject (..))
import Proarrow.Colimit.Pushout (HasPushouts (..))
import Proarrow.Core (CAT, CategoryOf (..), Is, Profunctor (..), Promonad (..), UN, dimapDefault)
import Proarrow.Limit.BinaryProduct
  ( HasBinaryProducts (..)
  , associatorProd
  , associatorProdInv
  , diag
  , leftUnitorProd
  , leftUnitorProdInv
  , rightUnitorProd
  , rightUnitorProdInv
  , swapProd
  )
import Proarrow.Limit.Equalizer (HasEqualizers (..))
import Proarrow.Limit.Pullback (HasPullbacks (..))
import Proarrow.Limit.Terminal (HasTerminalObject (..))
import Proarrow.Monoid (CocommutativeComonoid, Comonoid (..), Monoid (..))
import Proarrow.Optic (iso)
import Proarrow.Optic.Iso (Iso')
import Proarrow.Profunctor.Instance.Composition ((:.:) (..))

type data FINSET = FS Nat

type FinSet :: CAT FINSET
data FinSet a b where
  FinSet :: (SNatI n, SNatI m) => {unFinSet :: Vec n (Fin m)} -> FinSet (FS n) (FS m)

deriving instance P.Show (FinSet a b)
deriving instance P.Eq (FinSet a b)

instance Profunctor FinSet where
  dimap = dimapDefault
  r \\ FinSet{} = r
instance Promonad FinSet where
  id = FinSet universe
  FinSet l . FinSet r = FinSet (P.fmap (l !) r)

-- | The skeleton of the category of finite sets: objects are natural numbers and an arrow
-- @'FS' n '~>' 'FS' m@ is a function given by its table.
instance CategoryOf FINSET where
  type (~>) = FinSet
  type Ob a = (Is FS a, SNatI (UN FS a))

instance HasInitialObject FINSET where
  type InitialObject = FS Nat0
  initiate = FinSet VNil
instance HasBinaryCoproducts FINSET where
  type FS a || FS b = FS (Plus a b)
  withObCoprod @(FS a) @b r = case snat @a of
    SZ -> r
    SS @a' -> withObCoprod @_ @(FS a') @b r
  lft @(FS a) @(FS b) = withObCoprod @_ @(FS a) @(FS b) $ FinSet (P.fmap (weakenLeft (Proxy @b)) universe)
  rgt @(FS a) @(FS b) = withObCoprod @_ @(FS a) @(FS b) $ FinSet (P.fmap (weakenRight (Proxy @a)) universe)
  FinSet @a l ||| FinSet @b r = withObCoprod @_ @(FS a) @(FS b) $ FinSet (l ++ r)

instance HasTerminalObject FINSET where
  type TerminalObject = FS Nat1
  terminate = FinSet (repeat fin0)
instance HasBinaryProducts FINSET where
  type FS a && FS b = FS (Mult a b)
  withObProd @(FS a) @b r = case snat @a of
    SZ -> r
    SS @a' -> withObProd @_ @(FS a') @b $ withObCoprod @_ @b @(FS (Mult a' (UN FS b))) r
  fst @(FS a) @(FS b) = withObProd @_ @(FS a) @(FS b) $ FinSet (concat @a @b $ P.fmap repeat universe)
  snd @(FS a) @(FS b) = withObProd @_ @(FS a) @(FS b) $ FinSet (concat @a @b $ repeat universe)
  FinSet @_ @a l &&& FinSet @_ @b r = withObProd @_ @(FS a) @(FS b) $ FinSet (zipWith mult l r)

instance Distributive FINSET where
  distL @(FS a) @(FS b) @(FS c) =
    withObCoprod @_ @(FS b) @(FS c) $
      withObProd @_ @(FS a) @(FS (Plus b c)) $
        withObProd @_ @(FS a) @(FS b) $
          withObProd @_ @(FS a) @(FS c) $
            withObCoprod @_ @(FS (Mult a b)) @(FS (Mult a c)) $
              FinSet $
                concat @a @(Plus b c) $
                  P.fmap
                    ( \i ->
                        P.fmap (\j -> weakenLeft (Proxy @(Mult a c)) (mult @a @b i j)) (universe @b)
                          ++ P.fmap (\j -> weakenRight (Proxy @(Mult a b)) (mult @a @c i j)) (universe @c)
                    )
                    (universe @a)
  distR @(FS a) @(FS b) @(FS c) =
    withObCoprod @_ @(FS a) @(FS b) $
      withObProd @_ @(FS (Plus a b)) @(FS c) $
        withObProd @_ @(FS a) @(FS c) $
          withObProd @_ @(FS b) @(FS c) $
            withObCoprod @_ @(FS (Mult a c)) @(FS (Mult b c)) $
              FinSet $
                concat @(Plus a b) @c $
                  P.fmap (\i -> P.fmap (\j -> weakenLeft (Proxy @(Mult b c)) (mult @a @c i j)) (universe @c)) (universe @a)
                    ++ P.fmap (\i -> P.fmap (\j -> weakenRight (Proxy @(Mult a c)) (mult @b @c i j)) (universe @c)) (universe @b)
  absorbL @(FS a) = withObProd @_ @(FS a) @(FS Z) $ FinSet (concat @a @Z (repeat VNil))
  absorbR = FinSet VNil

-- | >>> import Data.Type.Nat
-- >>> import Data.Fin
-- >>> mult @Nat5 @Nat4 fin4 fin2 -- 4*4+2
-- 18
-- >>> mult @Nat4 @Nat5 fin2 fin4 -- 5*2+4
-- 14
mult :: forall n m. (SNatI n, SNatI m) => Fin n -> Fin m -> Fin (Mult n m)
mult n m = case snat @n of
  SZ -> case n of {}
  SS @n' -> case n of
    FZ -> weakenLeft (Proxy @(Mult n' m)) m
    FS n' -> weakenRight (Proxy @m) (mult n' m)

unmult :: forall n m. (SNatI n, SNatI m) => Fin (Mult n m) -> (Fin n, Fin m)
unmult f = case snat @n of
  SZ -> case f of {}
  SS @n' -> case split @m @(Mult n' m) f of
    P.Left m -> (FZ, m)
    P.Right f' -> let (n, m) = unmult @n' @m f' in (FS n, m)

instance MonoidalProfunctor FinSet where
  one = id
  (**) = (***)

instance Monoidal FINSET where
  type a ** b = a && b
  type Unit = FS Nat1
  withOb2 @a @b = withObProd @_ @a @b
  leftUnitor = leftUnitorProd
  leftUnitorInv = leftUnitorProdInv
  rightUnitor = rightUnitorProd
  rightUnitorInv = rightUnitorProdInv
  associator @a @b @c = associatorProd @a @b @c
  associatorInv @a @b @c = associatorProdInv @a @b @c

instance SymMonoidal FINSET where
  swap @a @b = swapProd @a @b

type family Exp (a :: Nat) (b :: Nat) :: Nat where
  Exp a Z = S Z
  Exp a (S n) = Mult a (Exp a n)

instance Closed FINSET where
  type FS a ~~> FS b = FS (Exp b a)
  withObExp @(FS a) @b r = case snat @a of
    SZ -> r
    SS @a' -> withObExp @_ @(FS a') @b $ withObProd @_ @b @(FS (Exp (UN FS b) a')) r
  curry @(FS a) @(FS b) (FinSet @_ @c f) = withObExp @_ @(FS b) @(FS c) $ FinSet (P.fmap exp (chunks @a @b f))
  apply @(FS a) @(FS b) =
    withObExp @_ @(FS a) @(FS b) $
      withObProd @_ @(FS (Exp b a)) @(FS a) $
        FinSet (concatMap @_ @a @_ @(Exp b a) unExp universe)

-- | >>> import Data.Type.Nat
-- >>> import Data.Fin
-- >>> exp @_ @Nat2 (fin1 ::: fin0 ::: fin1 ::: fin1 ::: VNil)
-- 11
exp :: forall n m. (SNatI n, SNatI m) => Vec n (Fin m) -> Fin (Exp m n)
exp VNil = FZ
exp (x ::: xs) = case snat @n of
  SS @n' -> withObExp @_ @(FS n') @(FS m) $ mult x (exp @n' @m xs)

-- | >>> import Data.Type.Nat
-- >>> import Data.Fin
-- >>> unExp @Nat3 @Nat2 fin6
-- 1 ::: 1 ::: 0 ::: VNil
unExp :: forall n m. (SNatI n, SNatI m) => Fin (Exp m n) -> Vec n (Fin m)
unExp f = case snat @n of
  SZ -> VNil
  SS @n' -> withObExp @_ @(FS n') @(FS m) $ let (x, xs) = unmult @m @(Exp m n') f in x ::: unExp @n' @m xs

-- | >>> import Data.Type.Nat
-- >>> comult @(FS Nat4)
-- FinSet {unFinSet = 0 ::: 5 ::: 10 ::: 15 ::: VNil}
instance (SNatI a) => Comonoid (FS a) where
  counit = terminate
  comult = diag

instance (SNatI a) => CocommutativeComonoid (FS a)

instance CopyDiscard FINSET

instance Monoid (FS Nat1) where
  mempty = terminate
  mappend = terminate

-- | Finds an isomorphism between 'FS n' and itself that's consistent with the given (source, target)
-- pairs, if one exists.
findIso :: forall n. (SNatI n) => [(Fin n, Fin n)] -> P.Maybe (Iso' (FS n) (FS n))
findIso ps = mkIso P.<$> findBijection ps
  where
    mkIso :: Vec n (Fin n) -> Iso' (FS n) (FS n)
    mkIso fwd = iso (FinSet fwd) (FinSet (tabulate (\j -> findIndex (P.== j) fwd)))

-- | Extends the given (source, target) pairs to a full bijection on @Fin n@, if they're consistent
-- with being a partial injection (checked in both directions as they're added, so two different
-- sources claiming the same target is rejected just as readily as one source getting conflicting
-- targets). Unconstrained sources are matched up with whatever targets are left over, in order.
findBijection :: forall n. (SNatI n) => [(Fin n, Fin n)] -> P.Maybe (Vec n (Fin n))
findBijection ps = do
  (fwd, bwd) <- go (repeat P.Nothing) (repeat P.Nothing) ps
  let freeSrcs = P.filter (\i -> isNothing (fwd ! i)) (toList universe)
      freeTgts = P.filter (\j -> isNothing (bwd ! j)) (toList universe)
      completion = P.zip freeSrcs freeTgts
  P.pure (tabulate (\i -> fromMaybe (fromJust (List.lookup i completion)) (fwd ! i)))
  where
    go
      :: Vec n (P.Maybe (Fin n))
      -> Vec n (P.Maybe (Fin n))
      -> [(Fin n, Fin n)]
      -> P.Maybe (Vec n (P.Maybe (Fin n)), Vec n (P.Maybe (Fin n)))
    go fwd bwd [] = P.Just (fwd, bwd)
    go fwd bwd ((s, t) : rest) = case (fwd ! s, bwd ! t) of
      (P.Just t', _) | t' P./= t -> P.Nothing
      (_, P.Just s') | s' P./= s -> P.Nothing
      _ ->
        go
          (tabulate (\i -> if i P.== s then P.Just t else fwd ! i))
          (tabulate (\j -> if j P.== t then P.Just s else bwd ! j))
          rest

-- | >>> import Data.Fin
-- >>> import Data.Type.Nat
-- >>> import Data.Vec.Lazy
-- >>> let f :: FinSet (FS Nat4) (FS Nat3) = FinSet $ fin0 ::: fin1 ::: fin1 ::: fin0 ::: VNil
-- >>> let g :: FinSet (FS Nat4) (FS Nat3) = FinSet $ fin2 ::: fin0 ::: fin1 ::: fin0 ::: VNil
-- >>> let h :: FinSet (FS Nat3) (FS Nat4) = FinSet $ fin3 ::: fin2 ::: fin3 ::: VNil
-- >>> (equalize f g \incl -> let p = factorEqualizer incl h in P.show (incl, p, incl . p)) :: P.String
-- "(FinSet {unFinSet = 2 ::: 3 ::: VNil},FinSet {unFinSet = 1 ::: 0 ::: 1 ::: VNil},FinSet {unFinSet = 3 ::: 2 ::: 3 ::: VNil})"
instance HasEqualizers FINSET where
  equalize (FinSet f) (FinSet g) k =
    let groups = [x | x <- toList universe, f ! x P.== g ! x]
    in reifyList groups \vec -> k (FinSet vec)
  factorEqualizer (FinSet incl) (FinSet h) = FinSet (tabulate (\c -> findIndex (P.== (h ! c)) incl))

-- Example 3.84 of Seven Sketches (A: 0=red, 1=blue, 2=black)

-- | >>> import Data.Fin
-- >>> import Data.Type.Nat
-- >>> import Data.Vec.Lazy
-- >>> let f :: FinSet (FS Nat6) (FS Nat3) = FinSet $ fin0 ::: fin1 ::: fin0 ::: fin0 ::: fin2 ::: fin1 ::: VNil
-- >>> let g :: FinSet (FS Nat4) (FS Nat3) = FinSet $ fin2 ::: fin0 ::: fin1 ::: fin0 ::: VNil
-- >>> (pullback f g \(FinSet l) (FinSet r) -> P.show (l, r)) :: P.String
-- "(0 ::: 0 ::: 1 ::: 2 ::: 2 ::: 3 ::: 3 ::: 4 ::: 5 ::: VNil,1 ::: 3 ::: 2 ::: 1 ::: 3 ::: 1 ::: 3 ::: 0 ::: 2 ::: VNil)"
instance HasPullbacks FINSET where
  pullback (FinSet f) (FinSet g) k =
    let
      gByValue = IM.fromListWith (P.flip (P.++)) [(P.fromEnum (g ! y), [y]) | y <- toList universe]
      groups = [(x, y) | x <- toList universe, y <- fromMaybe [] (IM.lookup (P.fromEnum (f ! x)) gByValue)]
    in
      reifyList groups \vec -> k (FinSet $ P.fmap fst vec) (FinSet $ P.fmap snd vec)

instance HasCoequalizers FINSET where
  coequalize (FinSet @_ @a f) (FinSet g) k =
    let
      find m i = P.maybe i (find m) $ IM.lookup (P.fromEnum i) m
      union m (i, j) = let ri = find m i; rj = find m j in if ri P.== rj then m else IM.insert (P.fromEnum ri) rj m
      unionFind = P.foldl union IM.empty (zipWith (,) f g)
      step m x = IM.insertWith (P.++) (P.fromEnum $ find unionFind x) [x] m
      groups = IM.elems $ P.foldl step IM.empty (universe @a)
    in
      reifyList groups \vec -> k (FinSet (tabulate (\a -> findIndex (P.elem a) vec)))
  factorCoequalizer (FinSet q) (FinSet h) =
    let reps = IM.fromListWith (\_ old -> old) [(P.fromEnum (q ! b), b) | b <- toList universe]
    in FinSet (tabulate (\i -> h ! (reps IM.! P.fromEnum i)))

-- Exercise 6.22 of Seven Sketches

-- | >>> import Data.Fin
-- >>> import Data.Type.Nat
-- >>> let l :: FinSet (FS Nat4) (FS Nat3) = FinSet $ fin0 ::: fin0 ::: fin1 ::: fin2 ::: VNil
-- >>> let r :: FinSet (FS Nat4) (FS Nat5) = FinSet $ fin0 ::: fin2 ::: fin4 ::: fin4 ::: VNil
-- >>> (pushout l r \(FinSet l') (FinSet r') -> P.show (l', r')) :: P.String
-- "(1 ::: 3 ::: 3 ::: VNil,1 ::: 0 ::: 1 ::: 2 ::: 3 ::: VNil)"
instance HasPushouts FINSET

findIndex :: (a -> P.Bool) -> Vec n a -> Fin n
findIndex _ VNil = P.error "unexpected missing element"
findIndex f (a ::: as)
  | f a = FZ
  | P.otherwise = FS $ findIndex f as

-- | >>> import Proarrow.Colimit.Pushout (isEpi)
-- >>> import Data.Fin
-- >>> import Data.Type.Nat
-- >>> let f :: FinSet (FS Nat3) (FS Nat3) = FinSet $ fin2 ::: fin0 ::: fin1 ::: VNil
-- >>> (pushout f f \(FinSet g1) (FinSet g2) -> P.show (g1, g2)) :: P.String
-- "(0 ::: 1 ::: 2 ::: VNil,0 ::: 1 ::: 2 ::: VNil)"
-- >>> isEpi f
-- True
-- >>> import Proarrow.Limit.Pullback (isMono)
-- >>> (pullback f f \(FinSet l) (FinSet r) -> P.show (l, r)) :: P.String
-- "(0 ::: 1 ::: 2 ::: VNil,0 ::: 1 ::: 2 ::: VNil)"
-- >>> isMono f
-- True
-- >>> import Proarrow.Category.Topos (classifyImage, classifyKernelPair, and, or, implies, false)
-- >>> (classifyImage f, classifyKernelPair f)
-- (FinSet {unFinSet = 1 ::: 1 ::: 1 ::: VNil},FinSet {unFinSet = 1 ::: 0 ::: 0 ::: 0 ::: 1 ::: 0 ::: 0 ::: 0 ::: 1 ::: VNil})
-- >>> [and, or, implies] :: [FinSet (FS Nat4) (FS Nat2)]
-- [FinSet {unFinSet = 0 ::: 0 ::: 0 ::: 1 ::: VNil},FinSet {unFinSet = 0 ::: 1 ::: 1 ::: 1 ::: VNil},FinSet {unFinSet = 1 ::: 1 ::: 0 ::: 1 ::: VNil}]
-- >>> false :: FinSet (FS Nat1) (FS Nat2)
-- FinSet {unFinSet = 0 ::: VNil}
instance HasSubobjectClassifier FINSET where
  type Omega = FS Nat2
  true = FinSet $ fin1 ::: VNil
  classifyGraph (FinSet @n @m f) = withObProd @_ @(FS n) @(FS m) $ FinSet $ tabulate
    \(unmult @n @m -> (n, m)) -> if f ! n P.== m then fin1 else fin0

instance HasEpiMonoFactorization FINSET where
  factorize (FinSet f) = reifyList (nubOrd (toList f)) \vec ->
    let revMap = IM.fromList (toList (zipWith (\k v -> (P.fromEnum k, v)) vec universe))
    in FinSet (tabulate (\a -> revMap IM.! P.fromEnum (f ! a)))
         :.: FinSet vec

instance ElementaryTopos FINSET