{-# LANGUAGE DataKinds #-}
{-# LANGUAGE FlexibleInstances #-}
{-# LANGUAGE MultiParamTypeClasses #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE TypeApplications #-}
{-# LANGUAGE TypeFamilies #-}
{-# LANGUAGE UndecidableInstances #-}
{-# OPTIONS_GHC -Wno-orphans #-}
module Circuit.Mat.Array
(
IxMap,
ArrayC (..),
LiftFun (..),
Tabular (..),
reindex,
batch,
tabulateC,
indexC,
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, (.))
type IxMap (s' :: [Nat]) (s :: [Nat]) = Fins s' -> Fins s
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
}
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
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"
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 ::
(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)
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)
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
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))