proarrow-0.3.0.0: src/Proarrow/Tools/Einsum.hs
{-# LANGUAGE AllowAmbiguousTypes #-}
-- | Einstein summation with a numpy-style specification, as in @'einsum' \@"ij,jk->ik" a b@, in any
-- hypergraph category. A tensor is a state, @'Tensor' xs@, whose type lists the objects of its
-- indices. Each letter of the specification is an index; the letters of the inputs are matched with
-- the objects of the tensors given, so a letter used with two different objects does not compile,
-- and the result's objects are those of the output letters. Without @->@ the output is, as in
-- numpy, the letters used once, in alphabetical order. A letter may also occur more than once in the
-- output, which copies it.
--
-- The specification is an open hypergraph ("Proarrow.Category.Instance.OpenHypergraph"): a node for
-- each letter, the tensors as boxes, and the output letters as its boundary. It is already in normal
-- form, and the result is its 'Proarrow.Category.Instance.OpenHypergraph.simplify': the tensors
-- contracted in pairs, each letter summed out by a spider as soon as nothing still to come has it.
module Proarrow.Tools.Einsum
( Tensor
, einsum
, Einsum
, EinsumType
, Inputs
, Output
) where
import Data.Containers.ListUtils (nubOrd)
import Data.Kind (Constraint, Type)
import Data.Map.Strict qualified as M
import Data.Proxy (Proxy (..))
import Data.Type.Bool (type (&&), type (||))
import Data.Type.Equality (type (==))
import GHC.TypeLits (CmpChar, ErrorMessage (..), KnownChar, Symbol, TypeError, UnconsSymbol, charVal)
import GHC.TypeNats (Nat, type (+))
import Prelude (Bool (..), Char, Maybe (..), Ordering (..), type (~))
import Prelude qualified as P
import Proarrow.Category.Instance.OpenHypergraph
( Box (..)
, SIMPLIFY
, SomeArrow (..)
, SortList
, Wires
, simplify
, someArrow
, unsafeOpenHypergraph
, unsafePrim
)
import Proarrow.Category.Instance.Product (Fst, Snd)
import Proarrow.Category.Monoidal (State)
import Proarrow.Category.Monoidal.Hypergraph (Hypergraph, Sized)
import Proarrow.Category.Monoidal.Strictified (Fold, type (++))
import Proarrow.Core (CategoryOf (..), Kind)
import Proarrow.Functor (FunctorForRep (..))
import Proarrow.Object (KnownListOf (..), mapListOf, someOfList)
-- | A tensor with indices of the given objects: a state of their tensor, as a morphism of
-- 'Proarrow.Category.Monoidal.Strictified.Strictified' from @'[]@.
type Tensor :: forall k. [k] -> Type
type Tensor xs = State xs
-- The specification
-- | A specification parsed into the letters of each input and of the output.
type Parse :: Symbol -> ([[Char]], [Char])
type Parse s = ParseInputs (UnconsSymbol s) '[] '[]
-- | The letters of each input.
type Inputs :: Symbol -> [[Char]]
type Inputs s = Fst @ Parse s
-- | The letters of the output.
type Output :: Symbol -> [Char]
type Output s = Snd @ Parse s
-- the letters of the current input, reversed, and the inputs before it, reversed
type ParseInputs :: Maybe (Char, Symbol) -> [Char] -> [[Char]] -> ([[Char]], [Char])
type family ParseInputs m cur acc where
ParseInputs 'Nothing cur acc = Implicit (Finish cur acc)
ParseInputs ('Just '( ',', s)) cur acc = ParseInputs (UnconsSymbol s) '[] (Reverse cur ': acc)
ParseInputs ('Just '( ' ', s)) cur acc = ParseInputs (UnconsSymbol s) cur acc
ParseInputs ('Just '( '-', s)) cur acc = ParseArrow (UnconsSymbol s) (Finish cur acc)
ParseInputs ('Just '(c, s)) cur acc = ParseInputs (UnconsSymbol s) (c ': cur) acc
-- the inputs, with the letters of the last one
type Finish :: [Char] -> [[Char]] -> [[Char]]
type Finish cur acc = Reverse (Reverse cur ': acc)
type ParseArrow :: Maybe (Char, Symbol) -> [[Char]] -> ([[Char]], [Char])
type family ParseArrow m ins where
ParseArrow ('Just '( '>', s)) ins = '(ins, ParseOutput (UnconsSymbol s) '[])
ParseArrow _ _ = TypeError (Text "Proarrow.Tools.Einsum: expected > after - in the specification")
-- | The inputs with numpy's implicit output: the letters used once, in alphabetical order.
type Implicit :: [[Char]] -> ([[Char]], [Char])
type Implicit ins = '(ins, Sort (Once (Fold ins) (Fold ins)))
-- the letters of the first list that occur once in the second
type Once :: [Char] -> [Char] -> [Char]
type family Once cs all where
Once '[] all = '[]
Once (c ': cs) all = OnceIf (Count c all) c (Once cs all)
type OnceIf :: Nat -> Char -> [Char] -> [Char]
type family OnceIf n c cs where
OnceIf 1 c cs = c ': cs
OnceIf _ c cs = cs
type Count :: Char -> [Char] -> Nat
type family Count c cs where
Count c '[] = 0
Count c (c ': cs) = 1 + Count c cs
Count c (d ': cs) = Count c cs
type Sort :: [Char] -> [Char]
type family Sort cs where
Sort '[] = '[]
Sort (c ': cs) = InsertSorted c (Sort cs)
type InsertSorted :: Char -> [Char] -> [Char]
type family InsertSorted c cs where
InsertSorted c '[] = '[c]
InsertSorted c (d ': ds) = InsertOrd (CmpChar c d) c d ds
type InsertOrd :: Ordering -> Char -> Char -> [Char] -> [Char]
type family InsertOrd o c d ds where
InsertOrd 'GT c d ds = d ': InsertSorted c ds
InsertOrd _ c d ds = c ': d ': ds
type ParseOutput :: Maybe (Char, Symbol) -> [Char] -> [Char]
type family ParseOutput m cur where
ParseOutput 'Nothing cur = Reverse cur
ParseOutput ('Just '( ' ', s)) cur = ParseOutput (UnconsSymbol s) cur
ParseOutput ('Just '(c, s)) cur = ParseOutput (UnconsSymbol s) (c ': cur)
type Reverse :: [a] -> [a]
type Reverse xs = ReverseOnto xs '[]
type ReverseOnto :: [a] -> [a] -> [a]
type family ReverseOnto xs acc where
ReverseOnto '[] acc = acc
ReverseOnto (x ': xs) acc = ReverseOnto xs (x ': acc)
-- The indices
-- | The objects of the indices, by letter, in the order the letters first appear.
type Env :: Kind -> Kind
type Env k = [(Char, k)]
-- | The letters of the inputs matched with the objects of their tensors.
type BindAll :: forall k. [([Char], [k])] -> Env k -> Env k
type family BindAll ts env where
BindAll '[] env = env
BindAll ('(ls, xs) ': ts) env = BindAll ts (Bind ls xs env)
type Bind :: forall k. [Char] -> [k] -> Env k -> Env k
type family Bind ls xs env where
Bind (c ': ls) (x ': xs) env = Bind ls xs (Insert c x env)
Bind _ _ env = env
-- a letter already bound keeps its first object; 'Check' reports a second one
type Insert :: forall k. Char -> k -> Env k -> Env k
type family Insert c x env where
Insert c x '[] = '[ '(c, x)]
Insert c x ('(c, y) ': env) = '(c, y) ': env
Insert c x (p ': env) = p ': Insert c x env
-- stuck at a letter no input has, which 'Check' reports
type Lookup :: forall k. Char -> Env k -> k
type family Lookup c env where
Lookup c ('(c, x) ': env) = x
Lookup c (p ': env) = Lookup c env
-- | The errors of a specification, reported once: a tensor with a different number of indices than
-- letters, a character that is not a letter, and, when there is neither, a letter used with two
-- objects and an output letter no input has. The other constraints of 'Einsum' get stuck instead
-- of repeating them.
type Check :: forall k. [([Char], [k])] -> [Char] -> Env k -> Constraint
type Check ts out env = Letters (LettersOf ts ++ out) (Arities ts (Agreements ts env, CheckOutput out env))
type LettersOf :: forall k. [([Char], [k])] -> [Char]
type family LettersOf ts where
LettersOf '[] = '[]
LettersOf ('(ls, xs) ': ts) = ls ++ LettersOf ts
-- the given constraint, when every index is a letter
type Letters :: [Char] -> Constraint -> Constraint
type family Letters ls c where
Letters '[] c = c
Letters (l ': ls) c = LetterIf (IsLetter l) l (Letters ls c)
type LetterIf :: Bool -> Char -> Constraint -> Constraint
type family LetterIf ok l c where
LetterIf 'True l c = c
LetterIf 'False l c =
TypeError (Text "Proarrow.Tools.Einsum: " :<>: ShowType l :<>: Text " is not a letter, so it cannot be an index")
type IsLetter :: Char -> Bool
type IsLetter c = Within 'a' c 'z' || Within 'A' c 'Z'
-- whether the middle character is between the outer two
type Within :: Char -> Char -> Char -> Bool
type Within lo c hi = NotGT (CmpChar lo c) && NotGT (CmpChar c hi)
type NotGT :: Ordering -> Bool
type family NotGT o where
NotGT 'GT = 'False
NotGT _ = 'True
-- the given constraint, when every tensor has as many letters as indices
type Arities :: forall k. [([Char], [k])] -> Constraint -> Constraint
type family Arities ts c where
Arities '[] c = c
Arities ('(ls, xs) ': ts) c = ArityError (Len ls == Len xs) ls xs (Arities ts c)
type ArityError :: forall k. Bool -> [Char] -> [k] -> Constraint -> Constraint
type family ArityError ok ls xs c where
ArityError 'True ls xs c = c
ArityError 'False ls xs c =
TypeError
( Text "Proarrow.Tools.Einsum: the tensor with indices "
:<>: ShowType xs
:<>: Text " has letters "
:<>: ShowType ls
)
type Agreements :: forall k. [([Char], [k])] -> Env k -> Constraint
type family Agreements ts env where
Agreements '[] env = ()
Agreements ('(c ': ls, x ': xs) ': ts) env = (Agrees c x env, Agreements ('(ls, xs) ': ts) env)
Agreements ('(ls, xs) ': ts) env = Agreements ts env
type Agrees :: forall k. Char -> k -> Env k -> Constraint
type family Agrees c x env where
Agrees c x ('(c, x) ': env) = ()
Agrees c x ('(c, y) ': env) =
TypeError
( Text "Proarrow.Tools.Einsum: the index "
:<>: ShowType c
:<>: Text " is used with both "
:<>: ShowType y
:<>: Text " and "
:<>: ShowType x
)
Agrees c x (p ': env) = Agrees c x env
type CheckOutput :: forall k. [Char] -> Env k -> Constraint
type family CheckOutput out env where
CheckOutput '[] env = ()
CheckOutput (c ': out) env = (Bound c env, CheckOutput out env)
type Bound :: forall k. Char -> Env k -> Constraint
type family Bound c env where
Bound c ('(c, x) ': env) = ()
Bound c (p ': env) = Bound c env
Bound c '[] =
TypeError
(Text "Proarrow.Tools.Einsum: the output letter " :<>: ShowType c :<>: Text " is not an index of any input")
-- | The objects of the given letters.
type Objs :: forall k. [Char] -> Env k -> [k]
type family Objs ls env where
Objs '[] env = '[]
Objs (c ': ls) env = Lookup c env ': Objs ls env
type Len :: [a] -> Nat
type family Len xs where
Len '[] = 0
Len (x ': xs) = 1 + Len xs
-- The network
-- | The letters as a value.
type KnownChars :: [Char] -> Constraint
type KnownChars ls = KnownListOf KnownChar ls
chars :: forall ls. (KnownChars ls) => [Char]
chars = mapListOf @KnownChar (\ @c -> charVal (Proxy @c)) (listOf @KnownChar @ls)
-- | The open hypergraph of a specification: a node for each letter, of the sort of its object, a box
-- for each tensor with an output for each of its letters, and the output letters as the boundary.
-- The type checker has matched the letters with the objects of the tensors and of the output, so
-- the hypergraph needs no checks.
network
:: forall {k} (os :: [k])
. (SortList os)
=> [([Char], SomeArrow k)]
-> [Char]
-> Wires '[] ~> (Wires os :: SIMPLIFY k)
network tensors out =
unsafeOpenHypergraph
(P.fmap (sortOfLetter M.!) letters)
[]
(P.fmap (index M.!) out)
[Box (unsafePrim t) [] (P.fmap (index M.!) ls) | (ls, t) <- tensors]
where
-- the letters in the order they first appear, with their sorts
letters = nubOrd (P.concatMap P.fst tensors)
sortOfLetter = M.fromList [(l, x) | (ls, SomeArrow _ ys _) <- tensors, (l, x) <- P.zip ls (someOfList ys)]
index = M.fromList (P.zip letters [0 ..])
-- Einsum
-- | Collect the tensors of the inputs, then sum.
type Einsum :: forall {k}. [[Char]] -> [Char] -> [([Char], [k])] -> Type -> Constraint
class Einsum ins out (ts :: [([Char], [k])]) r where
collect :: [([Char], SomeArrow k)] -> r
instance
(r ~ (Tensor xs -> r'), KnownChars ls, SortList xs, Einsum ins out ('(ls, xs) ': ts) r')
=> Einsum (ls ': ins) out (ts :: [([Char], [k])]) r
where
collect acc t = collect @ins @out @('(ls, xs) ': ts) ((chars @ls, someArrow t) : acc)
-- The tensors are the boxes of an open hypergraph, which is read back with each box its tensor.
instance
( env ~ BindAll (Reverse ts) '[]
, Check ts out env
, Hypergraph k
, Sized k
, os ~ Objs out env
, r ~ Tensor os
, SortList os
, KnownChars out
)
=> Einsum '[] out (ts :: [([Char], [k])]) r
where
collect acc = simplify (network @os (P.reverse acc) (chars @out))
-- | Einstein summation: @einsum \@"ij,jk->ik" a b@ is the tensor with entries the sums over @j@ of
-- the products of the entries of @a@ and @b@. The tensors are given after the specification, one for
-- each input, and the result's type follows from theirs.
einsum :: forall {k} (s :: Symbol) r. (Einsum (Inputs s) (Output s) ('[] :: [([Char], [k])]) r) => r
einsum = collect @(Inputs s) @(Output s) @('[] :: [([Char], [k])]) []
-- | The type of @'einsum' \@s@ applied to tensors with indices of the given objects. Each letter
-- takes the object it is first given, so a letter given two objects is one type variable in every
-- input, as in
-- @'EinsumType' "ij,jk" '[ '[a, b], '[c, d]] = 'Tensor' '[a, b] -> 'Tensor' '[b, d] -> 'Tensor' '[a, d]@.
type EinsumType :: forall k. Symbol -> [[k]] -> Type
type EinsumType s xss = Arrows (Inputs s) (Output s) (BindAll (Zip (Inputs s) xss) '[])
type Zip :: [a] -> [b] -> [(a, b)]
type family Zip as bs where
Zip (a ': as) (b ': bs) = '(a, b) ': Zip as bs
Zip _ _ = '[]
type Arrows :: forall k. [[Char]] -> [Char] -> Env k -> Type
type family Arrows ins out env where
Arrows '[] out env = Tensor (Objs out env)
Arrows (ls ': ins) out env = Tensor (Objs ls env) -> Arrows ins out env