{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}
{-# OPTIONS_GHC -Wno-orphans #-}

-- | Arrays as the monoidal array category @[C; I]@ of Abbott & Zardini.
--
-- = Overview
--
-- Abbott & Zardini define an array category @[C; I]@ from any base symmetric
-- monoidal category @C@ and any indexing category @I@.  Objects are synthetic
-- arrays @[X; P]@; morphisms are generated by two families:
--
-- * __Batch lift__ @[f; P]@ — apply a base morphism @f : X -> Y@ at every
--   index of shape @P@.
-- * __Reindexing__ @[X; η]@ — precompose with an indexing morphism
--   @η : P -> Q@.
--
-- This module gives a direct presheaf semantics: an array of shape @s@ over
-- base category @arr@ is a morphism @Fins s -> X@ in @arr@.  The two
-- generators are therefore exactly categorical composition in @arr@.
--
-- = Instantiations
--
-- * @arr = (->)@: arrays are functions from indices to values.
-- * @arr = Mat r@: arrays are functional matrices between finite index sets.
--
-- The module also provides a conceptual renaming: @backpermute@ in harpie
-- is 'reindex' in @[C; I]@.
module Circuit.Mat.Array
  ( -- * Indexing category
    IxMap,

    -- * Array category carrier
    ArrayC (..),

    -- * Base-category support
    LiftFun (..),
    Tabular (..),

    -- * Array-category operations
    reindex,
    batch,
    tabulateC,
    indexC,

    -- * Finite support for shape indices
    allFins,
  )
where

import Circuit.Category (Category (..))
import Circuit.Mat (Finite (..), Mat (..), runMat)
import Data.Bool (bool)
import Data.Kind (Constraint, Type)
import GHC.TypeNats (Nat)
import Harpie.Shape (Fins (..), KnownNats (..), valuesOf)
import NumHask.Algebra.Additive (Additive (..))
import NumHask.Algebra.Multiplicative (Multiplicative (..))
import Prelude hiding (id, (.))

-- $setup
--
-- >>> :set -XDataKinds
-- >>> :set -XTypeApplications
-- >>> import Harpie.Shape
-- >>> import Circuit.Mat.Array

-- | A morphism in the indexing category @I@: a map between shape indices.
type IxMap (s' :: [Nat]) (s :: [Nat]) = Fins s' -> Fins s

-- | Arrays in the array category @[C; I]@.
--
-- An array of shape @s@ with values in object @a@ is a morphism
-- @Fins s -> a@ in the base category @arr@.
newtype ArrayC (arr :: Type -> Type -> Type) (s :: [Nat]) (a :: Type) = ArrayC
  { forall (arr :: * -> * -> *) (s :: [Nat]) a.
ArrayC arr s a -> arr (Fins s) a
runArrayC :: arr (Fins s) a
  }

-- | Lift a plain function to a morphism in the base category.
--
-- For functions this is the identity; for matrices it is the 'Fun'
-- constructor.
class (Category arr) => LiftFun arr where
  liftFun :: (i -> j) -> arr i j

instance LiftFun (->) where
  liftFun :: forall i j. (i -> j) -> i -> j
liftFun = (i -> j) -> i -> j
forall a. a -> a
forall k (arr :: k -> k -> *) (a :: k). Category arr => arr a a
id

instance LiftFun (Mat r) where
  liftFun :: forall i j. (i -> j) -> Mat r i j
liftFun = (i -> j) -> Mat r i j
forall i j s. (i -> j) -> Mat s i j
Fun

-- | Categories that can represent a functional relation @i -> j@ as a
-- morphism and recover it.
--
-- The @Mat r@ instance builds a matrix with a single 'one' in each row
-- and recovers the function by looking for that 'one'.  It therefore
-- requires the matrix to be functional.
class (Category arr) => Tabular arr where
  type TabulateOb arr (i :: Type) (j :: Type) :: Constraint

  tabulateArr :: (TabulateOb arr i j) => (i -> j) -> arr i j
  indexArr :: (TabulateOb arr i j) => arr i j -> (i -> j)

instance Tabular (->) where
  type TabulateOb (->) i j = ()
  tabulateArr :: forall i j. TabulateOb (->) i j => (i -> j) -> i -> j
tabulateArr = (i -> j) -> i -> j
forall a. a -> a
forall k (arr :: k -> k -> *) (a :: k). Category arr => arr a a
id
  indexArr :: forall i j. TabulateOb (->) i j => (i -> j) -> i -> j
indexArr = (i -> j) -> i -> j
forall a. a -> a
forall k (arr :: k -> k -> *) (a :: k). Category arr => arr a a
id

instance Tabular (Mat r) where
  type TabulateOb (Mat r) i j = (Finite i, Finite j, Eq j, Additive r, Multiplicative r, Eq r)
  tabulateArr :: forall i j. TabulateOb (Mat r) i j => (i -> j) -> Mat r i j
tabulateArr i -> j
f = (i -> j -> r) -> Mat r i j
forall i j s. (Finite i, Finite j) => (i -> j -> s) -> Mat s i j
Mat (\i
i j
j -> r -> r -> Bool -> r
forall a. a -> a -> Bool -> a
bool r
forall a. Additive a => a
zero r
forall a. Multiplicative a => a
one (i -> j
f i
i j -> j -> Bool
forall a. Eq a => a -> a -> Bool
== j
j))
  indexArr :: forall i j. TabulateOb (Mat r) i j => Mat r i j -> i -> j
indexArr Mat r i j
m i
i =
    case (j -> Bool) -> [j] -> [j]
forall a. (a -> Bool) -> [a] -> [a]
filter (\j
j -> Mat r i j -> i -> j -> r
forall s j i.
(Additive s, Multiplicative s, Eq j) =>
Mat s i j -> i -> j -> s
runMat Mat r i j
m i
i j
j r -> r -> Bool
forall a. Eq a => a -> a -> Bool
== r
forall a. Multiplicative a => a
one) [j]
forall a. Finite a => [a]
universe of
      [j
j] -> j
j
      [] -> [Char] -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Mat.Array.indexArr: no value at index"
      [j]
_ -> [Char] -> j
forall a. HasCallStack => [Char] -> a
error [Char]
"Circuit.Mat.Array.indexArr: non-functional matrix at index"

-- | Reindexing morphism @[X; η]@.
--
-- Precompose with an indexing morphism.  This is exactly harpie's
-- @backpermute@, renamed to match the paper.
reindex ::
  (LiftFun arr) =>
  IxMap s' s ->
  ArrayC arr s a ->
  ArrayC arr s' a
reindex :: forall (arr :: * -> * -> *) (s' :: [Nat]) (s :: [Nat]) a.
LiftFun arr =>
IxMap s' s -> ArrayC arr s a -> ArrayC arr s' a
reindex IxMap s' s
eta (ArrayC arr (Fins s) a
m) = arr (Fins s') a -> ArrayC arr s' a
forall (arr :: * -> * -> *) (s :: [Nat]) a.
arr (Fins s) a -> ArrayC arr s a
ArrayC (arr (Fins s) a
m arr (Fins s) a -> arr (Fins s') (Fins s) -> arr (Fins s') a
forall b c a. arr b c -> arr a b -> arr a c
forall k (arr :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category arr =>
arr b c -> arr a b -> arr a c
. IxMap s' s -> arr (Fins s') (Fins s)
forall i j. (i -> j) -> arr i j
forall (arr :: * -> * -> *) i j. LiftFun arr => (i -> j) -> arr i j
liftFun IxMap s' s
eta)

-- | Batch lift @[f; P]@.
--
-- Apply a base morphism at every index of the array.
batch ::
  (Category arr) =>
  arr a b ->
  ArrayC arr s a ->
  ArrayC arr s b
batch :: forall (arr :: * -> * -> *) a b (s :: [Nat]).
Category arr =>
arr a b -> ArrayC arr s a -> ArrayC arr s b
batch arr a b
f (ArrayC arr (Fins s) a
m) = arr (Fins s) b -> ArrayC arr s b
forall (arr :: * -> * -> *) (s :: [Nat]) a.
arr (Fins s) a -> ArrayC arr s a
ArrayC (arr a b
f arr a b -> arr (Fins s) a -> arr (Fins s) b
forall b c a. arr b c -> arr a b -> arr a c
forall k (arr :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category arr =>
arr b c -> arr a b -> arr a c
. arr (Fins s) a
m)

-- | Separator / tabulation: build an array from a function on indices.
--
-- This is the separator @Sp@ of the paper, mapping the presheaf
-- representation back to an array object.
tabulateC ::
  (Tabular arr, TabulateOb arr (Fins s) a) =>
  (Fins s -> a) ->
  ArrayC arr s a
tabulateC :: forall (arr :: * -> * -> *) (s :: [Nat]) a.
(Tabular arr, TabulateOb arr (Fins s) a) =>
(Fins s -> a) -> ArrayC arr s a
tabulateC Fins s -> a
f = arr (Fins s) a -> ArrayC arr s a
forall (arr :: * -> * -> *) (s :: [Nat]) a.
arr (Fins s) a -> ArrayC arr s a
ArrayC ((Fins s -> a) -> arr (Fins s) a
forall i j. TabulateOb arr i j => (i -> j) -> arr i j
forall (arr :: * -> * -> *) i j.
(Tabular arr, TabulateOb arr i j) =>
(i -> j) -> arr i j
tabulateArr Fins s -> a
f)

-- | Join / lookup: recover the function on indices from an array.
--
-- This is the join @Jn@ of the paper.
indexC ::
  (Tabular arr, TabulateOb arr (Fins s) a) =>
  ArrayC arr s a ->
  (Fins s -> a)
indexC :: forall (arr :: * -> * -> *) (s :: [Nat]) a.
(Tabular arr, TabulateOb arr (Fins s) a) =>
ArrayC arr s a -> Fins s -> a
indexC (ArrayC arr (Fins s) a
m) = arr (Fins s) a -> Fins s -> a
forall i j. TabulateOb arr i j => arr i j -> i -> j
forall (arr :: * -> * -> *) i j.
(Tabular arr, TabulateOb arr i j) =>
arr i j -> i -> j
indexArr arr (Fins s) a
m

-- | All indices of a statically-known shape.
allFins :: forall s. (KnownNats s) => [Fins s]
allFins :: forall (s :: [Nat]). KnownNats s => [Fins s]
allFins = [Int] -> Fins s
forall {k} (s :: k). [Int] -> Fins s
UnsafeFins ([Int] -> Fins s) -> [[Int]] -> [Fins s]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [Int] -> [[Int]]
forall {a}. (Num a, Enum a) => [a] -> [[a]]
enumerateShape (forall (s :: [Nat]). KnownNats s => [Int]
valuesOf @s)
  where
    enumerateShape :: [a] -> [[a]]
enumerateShape [] = [[]]
    enumerateShape (a
n : [a]
ns) = [a
i a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
rest | a
i <- [a
0 .. a
n a -> a -> a
forall a. Num a => a -> a -> a
- a
1], [a]
rest <- [a] -> [[a]]
enumerateShape [a]
ns]

instance (KnownNats s) => Finite (Fins s) where
  universe :: [Fins s]
universe = forall (s :: [Nat]). KnownNats s => [Fins s]
allFins @s

instance Functor (ArrayC (->) s) where
  fmap :: forall a b. (a -> b) -> ArrayC (->) s a -> ArrayC (->) s b
fmap a -> b
f (ArrayC Fins s -> a
g) = (Fins s -> b) -> ArrayC (->) s b
forall (arr :: * -> * -> *) (s :: [Nat]) a.
arr (Fins s) a -> ArrayC arr s a
ArrayC (a -> b
f (a -> b) -> (Fins s -> a) -> Fins s -> b
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall k (arr :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category arr =>
arr b c -> arr a b -> arr a c
. Fins s -> a
g)

instance Applicative (ArrayC (->) s) where
  pure :: forall a. a -> ArrayC (->) s a
pure a
a = (Fins s -> a) -> ArrayC (->) s a
forall (arr :: * -> * -> *) (s :: [Nat]) a.
arr (Fins s) a -> ArrayC arr s a
ArrayC (a -> Fins s -> a
forall a b. a -> b -> a
const a
a)
  ArrayC Fins s -> (a -> b)
f <*> :: forall a b.
ArrayC (->) s (a -> b) -> ArrayC (->) s a -> ArrayC (->) s b
<*> ArrayC Fins s -> a
x = (Fins s -> b) -> ArrayC (->) s b
forall (arr :: * -> * -> *) (s :: [Nat]) a.
arr (Fins s) a -> ArrayC arr s a
ArrayC (\Fins s
i -> Fins s -> (a -> b)
f Fins s
i (Fins s -> a
x Fins s
i))