proarrow-0.3.0.0: src/Proarrow/Category/Instance/TensorNetwork.hs
{-# LANGUAGE AllowAmbiguousTypes #-}
{-# LANGUAGE NoStarIsType #-}
-- | __Tensor networks__ as a hypergraph category: the same arrows as the matrices of
-- "Proarrow.Category.Instance.Mat", kept as a network instead of as one matrix. An object is a list
-- of dimensions, one for each wire. An arrow has nodes, each with a dimension, a node for each of its
-- input and output wires, and dense factors: flat vectors, each with a node for each of its axes. Its
-- entry at given indices of the wires is the sum, over the indices of the nodes that agree with
-- those of the wires, of the products of the entries of the factors.
--
-- So the structure of a hypergraph category costs nothing: identities, swaps, copying, merging,
-- discarding, cups and caps are wirings without factors, and the tensor puts factors side by side
-- without multiplying them. Composition glues the wirings and sums out the nodes that are no longer
-- on a wire, multiplying only the factors that a summed node joins. With "Proarrow.Tools.Einsum" a
-- network of tensors is contracted in the order its read-back chooses.
--
-- With the package's @blas@ flag, a contraction of two factors that is a matrix product is handed
-- to the system's BLAS (Accelerate on macOS, OpenBLAS elsewhere) at entries of type 'P.Double',
-- 'P.Float' and @'Complex' 'P.Double'@.
--
-- The biproducts are direct sums, a single wire whose dimension is the sum of the sizes of the two
-- objects. Their injections, projections and pairings are dense matrices, and so are the
-- distributors.
module Proarrow.Category.Instance.TensorNetwork
( TNET (..)
, TensorNetwork
, Scalar
, dimsOf
, Size
, fromVector
, toVector
, fromEntries
, entries
) where
import Data.Complex (Complex, conjugate)
import Data.Containers.ListUtils (nubOrd)
import Data.IntMap.Strict qualified as IM
import Data.IntSet qualified as IS
import Data.Kind (Constraint, Type)
import Data.List qualified as List
import Data.Ord (comparing)
import Data.Proxy (Proxy (..))
import Data.Set qualified as Set
import Data.Type.Equality ((:~:) (..))
import Data.Vector.Storable qualified as SV
import Data.Vector.Storable.Mutable qualified as MSV
import Data.Vector.Unboxed qualified as UV
import Foreign.Storable (Storable)
import GHC.TypeNats (KnownNat, Nat, SNat, natVal, sameNat, withKnownNat, withSomeSNat, type (*), type (+))
import Unsafe.Coerce (unsafeCoerce)
import Prelude (Int, ($), (*), (+), (-), (==))
import Prelude qualified as P
import Proarrow.Category.Enriched.Dagger (DaggerProfunctor (..))
import Proarrow.Category.Instance.FinHask (unionFind)
import Proarrow.Category.Instance.TensorNetwork.Blas (Gemm, gemmComplexDouble, gemmDouble, gemmFloat)
import Proarrow.Category.Monoidal (Monoidal (..), MonoidalProfunctor (..), SymMonoidal (..))
import Proarrow.Category.Monoidal.Action (MonoidalAction)
import Proarrow.Category.Monoidal.Closed (Closed (..))
import Proarrow.Category.Monoidal.CompactClosed (CompactClosed (..), coactCC)
import Proarrow.Category.Monoidal.CopyDiscard (CopyDiscard)
import Proarrow.Category.Monoidal.Dialogue (Dialogue (..))
import Proarrow.Category.Monoidal.Distributive (Distributive (..))
import Proarrow.Category.Monoidal.Hypergraph (Frobenius, Hypergraph, Sized (..), cap, cup)
import Proarrow.Category.Monoidal.IsoMix (IsoMix (..))
import Proarrow.Category.Monoidal.StarAutonomous (ExpSA, StarAutonomous (..), applySA, currySA, expSA)
import Proarrow.Category.Monoidal.Strength (Costrong (..))
import Proarrow.Colimit.BinaryCoproduct (HasBinaryCoproducts (..), HasBiproducts)
import Proarrow.Colimit.Initial (HasInitialObject (..))
import Proarrow.Core (CAT, CategoryOf (..), Is, Profunctor (..), Promonad (..), UN, dimapDefault, type (+->))
import Proarrow.Limit.BinaryProduct (HasBinaryProducts (..))
import Proarrow.Limit.Terminal (HasTerminalObject (..))
import Proarrow.Monoid (CocommutativeComonoid, CommutativeMonoid, Comonoid (..), Monoid (..))
import Proarrow.Object (KnownListOf (..), appendListOf, eqListOf, mapListOf, withKnownListOf, type (++))
import Proarrow.Optic.Iso (DecidableIso (..), isoFromEquality)
-- | The entries of the factors: numbers that can be stored unboxed. An instance needs no methods:
-- the loop that contracts factors is then compiled for its type.
type Scalar :: Type -> Constraint
class (P.Num e, Storable e) => Scalar e where
-- | The entries of a contraction of two factors: given the dimension of each axis of the result
-- and how far a step along it moves in the entries of each factor, the offset into each factor
-- of each summed index, and the entries of the two factors, each entry is a sum of products.
kernel :: [(Dim, Int, Int)] -> UV.Vector Index -> UV.Vector Index -> SV.Vector e -> SV.Vector e -> SV.Vector e
-- inlined into each instance, so that the loop is specialised to its type
kernel = contractLoop
{-# INLINE kernel #-}
-- | The conjugate of an entry, which the dagger takes: the identity on real numbers.
conj :: e -> e
conj = P.id
-- | A fast matrix product, which a contraction of two factors that is one is handed to.
gemm :: P.Maybe (Gemm e)
gemm = P.Nothing
instance Scalar Int
instance Scalar P.Double where
gemm = gemmDouble
instance Scalar P.Float where
gemm = gemmFloat
instance Scalar (Complex P.Double) where
conj = conjugate
gemm = gemmComplexDouble
-- | The loop that 'kernel' runs, written once for every type: a loop over each axis of the result,
-- the last one innermost, reading the entries of both factors with their steps, and for each
-- entry a loop over the summed indices.
{-# INLINEABLE contractLoop #-}
contractLoop
:: (P.Num e, Storable e)
=> [(Dim, Int, Int)] -> UV.Vector Index -> UV.Vector Index -> SV.Vector e -> SV.Vector e -> SV.Vector e
contractLoop steps !so !to !v !w = SV.create do
out <- MSV.unsafeNew (P.product [d | (d, _, _) <- steps])
let go [] !a !b !dst = MSV.unsafeWrite out dst (entry a b)
go [(!d, !s, !t)] !a !b !dst
| n == 1 = products 0 (a + so `UV.unsafeIndex` 0) (b + to `UV.unsafeIndex` 0)
| P.otherwise = sums 0 a b
where
-- with nothing summed, each entry is one product
products !i !x !y
| i == d = P.pure ()
| P.otherwise =
MSV.unsafeWrite out (dst + i) (v `SV.unsafeIndex` x * w `SV.unsafeIndex` y) P.>> products (i + 1) (x + s) (y + t)
sums !i !x !y
| i == d = P.pure ()
| P.otherwise = MSV.unsafeWrite out (dst + i) (entry x y) P.>> sums (i + 1) (x + s) (y + t)
go ((!d, !s, !t) : rest) !a !b !dst = outer 0 a b
where
size' = P.product [e | (e, _, _) <- rest]
outer !i !x !y
| i == d = P.pure ()
| P.otherwise = go rest x y (dst + i * size') P.>> outer (i + 1) (x + s) (y + t)
go steps 0 0 0
P.pure out
where
n = UV.length so
-- the sum of the products at a place in each factor; its arguments are strict, so that the loop
-- does not look them up again for every index
entry !a !b = sumFrom 0 0
where
sumFrom !i !acc =
if i == n
then acc
else
sumFrom (i + 1) (acc + v `SV.unsafeIndex` (a + so `UV.unsafeIndex` i) * w `SV.unsafeIndex` (b + to `UV.unsafeIndex` i))
-- | A node of a network, numbered from 0.
type Node = Int
-- | The dimension of a wire or a node.
type Dim = Int
-- | An index along a wire or an axis, or a position in a factor's entries.
type Index = Int
-- | Objects are lists of dimensions, one for each wire.
type data TNET (e :: Type) = TN [Nat]
-- | The dimensions of a list, as numbers.
dimsOf :: forall ns. (KnownListOf KnownNat ns) => [Dim]
dimsOf = mapListOf @KnownNat (\ @n -> P.fromIntegral (natVal (Proxy @n))) (listOf @KnownNat @ns)
-- | A dense factor: its entries, row by row over its axes, and the node of each axis.
data Factor e = Factor {axes :: [Node], values :: SV.Vector e}
-- | The nodes, numbered from 0 with their dimensions, the node of each input and output wire, and
-- the factors. Kept so that every node is on a wire and no factor has an axis twice.
data Net e = Net {nodeDims :: [Dim], ins :: [Node], outs :: [Node], factors :: [Factor e]}
-- | An arrow between two lists of dimensions.
type TensorNetwork :: CAT (TNET e)
data TensorNetwork a b where
TensorNetwork
:: forall {e} as bs. (KnownListOf KnownNat as, KnownListOf KnownNat bs) => Net e -> TensorNetwork (TN as :: TNET e) (TN bs)
-- | A wiring without factors: nodes of the given dimensions, and the node of each input and output
-- wire.
wiring
:: forall {e} as bs
. (KnownListOf KnownNat as, KnownListOf KnownNat bs)
=> [Dim] -> [Node] -> [Node] -> TensorNetwork (TN as :: TNET e) (TN bs)
wiring dims is os = TensorNetwork (Net dims is os [])
-- | The wires going straight through, between two lists of the same dimensions.
straight
:: forall {e} as bs. (KnownListOf KnownNat as, KnownListOf KnownNat bs) => TensorNetwork (TN as :: TNET e) (TN bs)
straight = wiring (dimsOf @as) (wires @as) (wires @as)
-- | The arrow with the given entries, laid out as 'toVector' gives them. The length of the vector
-- must be the product of all the dimensions.
fromVector
:: forall {e} as bs
. (Scalar e, KnownListOf KnownNat as, KnownListOf KnownNat bs)
=> SV.Vector e -> TensorNetwork (TN as :: TNET e) (TN bs)
fromVector vs
| SV.length vs P./= P.product dims = P.error "fromVector: the length is not the product of the dimensions"
| P.otherwise = TensorNetwork (Net dims is os [Factor (os P.++ is) vs])
where
dims = dimsOf @as P.++ dimsOf @bs
is = wires @as
os = [P.length is .. P.length dims - 1]
-- | The entries, row by row: a row for each index of the output wires and a column for each index
-- of the input wires, both with the last wire varying fastest. This relies on every node of the
-- network being on a wire, which composition keeps so.
toVector :: forall {e} a b. (Scalar e) => TensorNetwork (a :: TNET e) b -> SV.Vector e
toVector (TensorNetwork (Net dims is os fs))
| distinct ws = values (product ws)
| P.otherwise = SV.create do
out <- MSV.replicate (extent dims ws) 0
let vs = values (product order)
go [] !src !dst = MSV.unsafeWrite out dst (vs `SV.unsafeIndex` src) P.>> P.pure (src + 1)
go ((d, w) : rest) !src !dst = loop 0 src
where
loop !i !from
| i == d = P.pure from
| P.otherwise = go rest from (dst + i * w) P.>>= loop (i + 1)
_ <- go [(dims P.!! n, P.sum [st | (x, st) <- P.zip ws wireStrides, x == n]) | n <- order] 0 0
P.pure out
where
ws = os P.++ is
-- the product of the factors, kept over the given nodes
product keep = contract1 dims (case fs of [] -> unit; f : rest -> P.foldl times f rest) keep []
times g h = contract dims g h (nubOrd (axes g P.++ axes h)) []
-- with a node on several wires, each entry of the product over the nodes goes to the place
-- that its wires give it, and the others are 0
order = nubOrd ws
wireStrides = P.drop 1 (P.scanr (*) 1 (P.fmap (dims P.!!) ws))
-- | The arrow with the given rows of entries, laid out as 'entries' gives them.
fromEntries
:: forall {e} as bs
. (Scalar e, KnownListOf KnownNat as, KnownListOf KnownNat bs)
=> [[e]] -> TensorNetwork (TN as :: TNET e) (TN bs)
fromEntries rows = fromVector (SV.fromList (P.concat rows))
-- | The entries as rows, as 'toVector' lays them out.
entries :: forall {e} a b. (Scalar e) => TensorNetwork (a :: TNET e) b -> [[e]]
entries f@(TensorNetwork @as @bs _) = [SV.toList (SV.slice (r * ni) ni v) | r <- [0 .. size @bs - 1]]
where
v = toVector f
ni = size @as
-- | The matrix of a function from the indices of the inputs to those of the outputs: each column
-- has a 1 in the row the function gives it.
reindex
:: forall {e} as bs
. (Scalar e, KnownListOf KnownNat as, KnownListOf KnownNat bs)
=> (Index -> Index) -> TensorNetwork (TN as :: TNET e) (TN bs)
reindex f = fromVector (SV.generate (size @bs * ni) \k -> let (r, c) = k `P.quotRem` ni in if r == f c then 1 else 0)
where
ni = size @as
-- | The arrow with no entries, into or out of a wire of dimension 0.
zero
:: forall {e} as bs
. (Scalar e, KnownListOf KnownNat as, KnownListOf KnownNat bs)
=> TensorNetwork (TN as :: TNET e) (TN bs)
zero = fromVector SV.empty
-- | The size of a list of dimensions: their product.
type Size :: [Nat] -> Nat
type family Size ns where
Size '[] = 1
Size (n ': ns) = n * Size ns
-- | The size of a list of dimensions, as a number.
size :: forall ns. (KnownListOf KnownNat ns) => Int
size = P.product (dimsOf @ns)
-- | The dimension of the direct sum of two lists of dimensions.
withSum
:: forall as bs r. (KnownListOf KnownNat as, KnownListOf KnownNat bs) => ((KnownNat (Size as + Size bs)) => r) -> r
withSum = withKnownNat n
where
n :: SNat (Size as + Size bs)
n = withSomeSNat (P.fromIntegral (size @as + size @bs)) unsafeCoerce
-- | The product of two factors, with the given nodes kept, in that order, and the others summed.
contract :: forall e. (Scalar e) => [Dim] -> Factor e -> Factor e -> [Node] -> [Node] -> Factor e
contract dims f g keep summed
| P.Just mm <- gemm, P.Just r <- viaGemm mm dims f g keep summed = r
| P.otherwise =
Factor
keep
( kernel
[(dims P.!! x, stride dims f x, stride dims g x) | x <- keep]
(offsets f summed)
(offsets g summed)
(values f)
(values g)
)
where
-- the offset into a factor's entries of each index of the given nodes
offsets h = P.foldl (\acc x -> step acc (dims P.!! x) (stride dims h x)) (UV.singleton 0)
step acc d s = UV.generate (UV.length acc * d) \j -> let (q, r) = j `P.quotRem` d in acc `UV.unsafeIndex` q + r * s
-- | One factor with the given nodes kept, in that order, and the others summed: its product with the
-- unit. A factor with an axis twice gives its diagonal, and the entries are the same along a kept
-- node that it does not have.
contract1 :: (Scalar e) => [Dim] -> Factor e -> [Node] -> [Node] -> Factor e
contract1 dims f keep summed
| P.null summed P.&& axes f == keep = f
| P.otherwise = contract dims f unit keep summed
-- | The factor without axes whose entry is 1.
unit :: (P.Num e, Storable e) => Factor e
unit = Factor [] (SV.singleton 1)
-- | A contraction of two factors as a matrix product, when it is one: something is summed, every
-- summed node is on both factors, and the kept nodes of the first come before those of the second,
-- or after them, as the transposed product. A factor whose axes are in neither the order of its
-- matrix nor that of its transpose is rearranged first.
viaGemm :: (Scalar e) => Gemm e -> [Dim] -> Factor e -> Factor e -> [Node] -> [Node] -> P.Maybe (Factor e)
viaGemm mm dims f g keep summed
| P.not (matrixProduct f g summed) P.|| rows * cols * inner P.< gemmWork = P.Nothing
| keep == ms P.++ ns = P.Just (Factor keep (mm ta tb rows cols inner va vb))
| keep == ns P.++ ms = P.Just (Factor keep (mm (P.not tb) (P.not ta) cols rows inner vb va))
| P.otherwise = P.Nothing
where
ms = [x | x <- axes f, P.not (onFactor x g)]
ns = [x | x <- axes g, P.not (onFactor x f)]
ss = P.filter (`P.elem` summed) (axes f)
rows = extent dims ms
cols = extent dims ns
inner = extent dims ss
(ta, va) = matrix f ms ss
(tb, vb) = matrix g ss ns
-- a factor as a matrix with the first nodes as rows: transposed, or rearranged when it is not
-- already in that order
matrix h rs cs
| axes h == cs P.++ rs = (P.True, values h)
| P.otherwise = (P.False, values (contract1 dims h (rs P.++ cs) []))
-- | How far a step along a node moves in a factor's entries: the strides of its axes on that node,
-- 0 when it has none.
stride :: [Dim] -> Factor e -> Node -> Int
stride dims f node = P.sum [s | (a, s) <- P.zip (axes f) (P.drop 1 (P.scanr (*) 1 (P.fmap (dims P.!!) (axes f)))), a == node]
-- | Whether no element of the list is there twice.
distinct :: [Node] -> P.Bool
distinct = go IS.empty
where
go _ [] = P.True
go seen (x : xs) = P.not (x `IS.member` seen) P.&& go (IS.insert x seen) xs
-- | Whether the node is on one of the factor's axes.
onFactor :: Node -> Factor e -> P.Bool
onFactor x f = x `P.elem` axes f
-- | Whether a product of two factors that sums the given nodes is a matrix product: something is
-- summed, and the nodes on both factors are the summed ones.
matrixProduct :: Factor e -> Factor e -> [Node] -> P.Bool
matrixProduct f g summed = P.not (P.null summed) P.&& Set.fromList [x | x <- axes f, onFactor x g] == Set.fromList summed
-- | The number of indices of the given nodes together.
extent :: [Dim] -> [Node] -> Int
extent dims = P.product P.. P.fmap (dims P.!!)
-- | The number of multiplications from which a matrix product goes to 'gemm'.
gemmWork :: Int
gemmWork = 4096
-- | The network with every factor's axes distinct, every node that is on no wire summed out, and
-- the nodes renumbered.
normalize :: (Scalar e) => Net e -> Net e
normalize (Net dims is os fs) =
Net
(P.fmap (dims P.!!) kept)
(P.fmap renumber is)
(P.fmap renumber os)
(P.fmap relabel (scalars (loops P.++ sumOut reduced shared)))
where
summed = [n | n <- [0 .. P.length dims - 1], n `Set.notMember` keptSet]
-- the number of factors on each node
factorsOn n = IM.findWithDefault 0 n counts
distinctAxes = [(f, nubOrd (axes f)) | f <- fs]
counts = IM.fromListWith (+) [(m, 1 :: Int) | (_, as) <- distinctAxes, m <- as]
-- each factor with its axes distinct, and the summed nodes that no other factor has summed out
reduced =
[ let own = [m | m <- as, m `Set.notMember` keptSet, factorsOn m == 1] in contract1 dims f (as List.\\ own) own
| (f, as) <- distinctAxes
]
-- the summed nodes on no factor, as of closed loops: a scalar of their dimensions
loops = case [n | n <- summed, factorsOn n == 0] of
[] -> []
ns -> [Factor [] (SV.singleton (P.fromIntegral (extent dims ns)))]
shared = [n | n <- summed, factorsOn n P.> 1]
-- the factors on a shared node multiplied two at a time, first the pair whose result grows least
-- (its size less those of the two factors, as the read-back chooses), each node summed out by
-- the product that brings together the last two factors on it
sumOut gs [] = gs
sumOut gs ns@(n : _) =
let (touching, others) = List.partition (onFactor n) gs
numbered = P.zip [0 :: Int ..] touching
pairs = [(f, g, [h | (k, h) <- numbered, k P./= i, k P./= j] P.++ others) | (i, f) <- numbered, (j, g) <- numbered, i P.< j]
product (f, g, outside) =
let ds = [m | m <- ns, onFactor m f P.|| onFactor m g, P.not (P.any (onFactor m) outside)]
left = nubOrd (axes f P.++ axes g) List.\\ ds
-- a product that is not a matrix product puts the nodes that a later product sums
-- last, where a matrix product wants them
keep = if matrixProduct f g ds then left else let (later, now) = List.partition (`P.elem` ns) left in now P.++ later
in (extent dims keep - extent dims (axes f) - extent dims (axes g), (contract dims f g keep ds, outside, ds))
(_, (fg, rest, done)) = List.minimumBy (comparing P.fst) (P.fmap product pairs)
in sumOut (fg : rest) (ns List.\\ done)
-- the factors without axes multiplied into one
scalars gs = case List.partition (P.null P.. axes) gs of
(c : d : more, rest) -> Factor [] (SV.singleton (P.product [SV.head (values x) | x <- c : d : more])) : rest
_ -> gs
kept = nubOrd (is P.++ os)
keptSet = Set.fromList kept
renumber = renumbering kept
relabel (Factor as vs) = Factor (P.fmap renumber as) vs
-- | The new number of each of the given nodes: its place in the list.
renumbering :: [Node] -> Node -> Node
renumbering ns = (IM.fromList (P.zip ns [0 ..]) IM.!)
-- | Two networks side by side.
besideNet :: Net e -> Net e -> Net e
besideNet (Net d1 i1 o1 f1) (Net d2 i2 o2 f2) =
Net (d1 P.++ d2) (i1 P.++ P.fmap (+ k) i2) (o1 P.++ P.fmap (+ k) o2) (f1 P.++ P.fmap shift f2)
where
k = P.length d1
shift (Factor as vs) = Factor (P.fmap (+ k) as) vs
-- | The second network after the first: their wirings glued along the wires in between, and the
-- nodes that are then on no wire summed out.
composeNet :: (Scalar e) => Net e -> Net e -> Net e
composeNet f g
| P.Just through <- permutation g = f{outs = P.fmap (outs f P.!!) through}
| P.Just through <- permutation (transposeNet f) = g{ins = P.fmap (ins g P.!!) through}
| P.otherwise =
normalize
( Net
[nodeDims both P.!! r | r <- live]
(P.fmap node (ins f))
(P.fmap (node P.. (+ k)) (outs g))
[Factor (P.fmap node as) vs | Factor as vs <- factors both]
)
where
both = besideNet f g
k = P.length (nodeDims f)
find = unionFind (P.zip (outs f) (P.fmap (+ k) (ins g)))
-- the nodes that are left after gluing, numbered from 0
live = nubOrd [find n | n <- [0 .. P.length (nodeDims both) - 1]]
node = renumbering live P.. find
-- | For a wiring without factors whose every node has one input wire and at least one output wire,
-- such as a permutation or copying: the input wire that each output wire continues.
permutation :: Net e -> P.Maybe [Int]
permutation (Net _ is os fs)
| P.not (P.null fs) P.|| P.not (distinct is) = P.Nothing
| is == os = P.Just [0 .. P.length is - 1]
| Set.fromList is == Set.fromList os = P.Just (P.fmap (renumbering is) os)
| P.otherwise = P.Nothing
-- | The network the other way round.
transposeNet :: Net e -> Net e
transposeNet (Net d i o fs) = Net d o i fs
withAppend
:: forall as bs r. (KnownListOf KnownNat as, KnownListOf KnownNat bs) => ((KnownListOf KnownNat (as ++ bs)) => r) -> r
withAppend = withKnownListOf (appendListOf (listOf @KnownNat @as) (listOf @KnownNat @bs))
withAssoc
:: forall as bs cs r
. (KnownListOf KnownNat as, KnownListOf KnownNat bs, KnownListOf KnownNat cs)
=> ((KnownListOf KnownNat ((as ++ bs) ++ cs), KnownListOf KnownNat (as ++ (bs ++ cs))) => r) -> r
withAssoc r = withAppend @as @bs (withAppend @(as ++ bs) @cs (withAppend @bs @cs (withAppend @as @(bs ++ cs) r)))
-- | The arrow the other way round: its wiring with inputs and outputs swapped.
transpose :: TensorNetwork (a :: TNET e) b -> TensorNetwork b a
transpose (TensorNetwork n) = TensorNetwork (transposeNet n)
instance (Scalar e) => Profunctor (TensorNetwork :: CAT (TNET e)) where
dimap = dimapDefault
r \\ TensorNetwork{} = r
instance (Scalar e) => Promonad (TensorNetwork :: CAT (TNET e)) where
id = straight
TensorNetwork g . TensorNetwork f = TensorNetwork (composeNet f g)
-- | Tensor networks with entries @e@, between lists of dimensions.
instance (Scalar e) => CategoryOf (TNET e) where
type (~>) = TensorNetwork
type Ob a = (Is TN a, KnownListOf KnownNat (UN TN a))
-- | The conjugate transpose. 'dual' is the transpose without conjugating, since the compact-closed
-- structure is bilinear.
instance (Scalar e) => DaggerProfunctor (TensorNetwork :: CAT (TNET e)) where
dagger (TensorNetwork n) = TensorNetwork (transposeNet n){factors = [Factor as (SV.map conj vs) | Factor as vs <- factors n]}
-- | A wire of dimension 0.
instance (Scalar e) => HasInitialObject (TNET e) where
type InitialObject = TN '[0]
initiate = zero
-- | A wire of dimension 0.
instance (Scalar e) => HasTerminalObject (TNET e) where
type TerminalObject = TN '[0]
terminate = zero
-- | The direct sum, as a single wire: the transpose of the product.
instance (Scalar e) => HasBinaryCoproducts (TNET e) where
type a || b = TN '[Size (UN TN a) + Size (UN TN b)]
withObCoprod @(TN as) @(TN bs) r = withSum @as @bs r
lft @(TN as) @(TN bs) = withSum @as @bs (reindex P.id)
rgt @(TN as) @(TN bs) = withSum @as @bs (reindex (size @as +))
f ||| g = transpose (transpose f &&& transpose g)
-- | The direct sum, as a single wire: the entries of the first arrow above those of the second.
instance (Scalar e) => HasBinaryProducts (TNET e) where
type a && b = TN '[Size (UN TN a) + Size (UN TN b)]
withObProd @(TN as) @(TN bs) r = withSum @as @bs r
fst @a @b = transpose (lft @_ @a @b)
snd @a @b = transpose (rgt @_ @a @b)
f@(TensorNetwork @_ @as _) &&& g@(TensorNetwork @_ @bs _) = withSum @as @bs (fromVector (toVector f SV.++ toVector g))
instance (Scalar e) => HasBiproducts (TNET e)
instance (Scalar e) => MonoidalProfunctor (TensorNetwork :: CAT (TNET e)) where
one = id
TensorNetwork @as @bs f ** TensorNetwork @cs @ds g = withAppend @as @cs (withAppend @bs @ds (TensorNetwork (besideNet f g)))
-- | The wires side by side as the tensor: on matrices, the Kronecker product.
instance (Scalar e) => Monoidal (TNET e) where
type Unit = TN '[]
type a ** b = TN (UN TN a ++ UN TN b)
withOb2 @(TN as) @(TN bs) r = withAppend @as @bs r
leftUnitor = id
leftUnitorInv = id
rightUnitor @(TN as) = withAppend @as @'[] straight
rightUnitorInv @(TN as) = withAppend @as @'[] straight
associator @(TN as) @(TN bs) @(TN cs) = withAssoc @as @bs @cs straight
associatorInv @(TN as) @(TN bs) @(TN cs) = withAssoc @as @bs @cs straight
instance (Scalar e) => SymMonoidal (TNET e) where
swap @(TN as) @(TN bs) =
withAppend @as @bs $
withAppend @bs @as $
let na = P.length (dimsOf @as)
nb = P.length (dimsOf @bs)
in wiring (dimsOf @as P.++ dimsOf @bs) [0 .. na P.+ nb - 1] ([na .. na P.+ nb - 1] P.++ [0 .. na - 1])
-- | The distributors reorder the entries of the direct sums.
instance (Scalar e) => Distributive (TNET e) where
distL @(TN as) @(TN bs) @(TN cs) =
withSum @bs @cs $
withAppend @as @'[Size bs + Size cs] $
withAppend @as @bs $
withAppend @as @cs $
withSum @(as ++ bs) @(as ++ cs) $
let (sa, sb, sc) = (size @as, size @bs, size @cs)
target c = let (i, k) = c `P.quotRem` (sb + sc) in if k P.< sb then i * sb + k else sa * sb + i * sc + k - sb
in reindex target
distR @(TN as) @(TN bs) @(TN cs) =
withSum @as @bs $
withAppend @as @cs $
withAppend @bs @cs $
withSum @(as ++ cs) @(bs ++ cs) $
reindex P.id
absorbL @(TN as) = withAppend @as @'[0] zero
absorbR = zero
-- | Every object is self-dual, with the transpose as the dual of an arrow: its wiring the other way
-- round.
instance (Scalar e) => Dialogue (TNET e) where
type Dual a = a
withObDual r = r
dual = transpose
linDist @(TN as) @(TN bs) @(TN cs) (TensorNetwork (Net d i o f)) =
withAppend @bs @cs $ let na = P.length (dimsOf @as) in TensorNetwork (Net d (P.take na i) (P.drop na i P.++ o) f)
linDistInv @(TN as) @(TN bs) (TensorNetwork (Net d i o f)) =
withAppend @as @bs $ let nb = P.length (dimsOf @bs) in TensorNetwork (Net d (i P.++ P.take nb o) (P.drop nb o) f)
doubleNegInv = id
instance (Scalar e) => StarAutonomous (TNET e) where
dualInv = transpose
doubleNeg = id
instance (Scalar e) => Closed (TNET e) where
type a ~~> b = ExpSA a b
withObExp @(TN as) @(TN bs) r = withAppend @as @bs r
curry @a @b = currySA @a @b
apply @a @b = applySA @a @b
(^^^) = expSA
instance (Scalar e) => IsoMix (TNET e) where
dualUnit = id
dualUnitInv = id
dualityCounit @a = cap @a
instance (Scalar e) => CompactClosed (TNET e) where
distribDual @(TN as) @(TN bs) = withAppend @as @bs id
dualityUnit @a = cup @a
instance (Scalar e, MonoidalAction (t :: (TNET e, TNET e) +-> TNET e)) => Costrong t (TensorNetwork :: CAT (TNET e)) where
coact @x = coactCC @t @x
-- | Merging each wire with its partner, and the unit.
instance (Scalar e, KnownListOf KnownNat ns) => Monoid (TN ns :: TNET e) where
mempty = wiring (dimsOf @ns) [] (wires @ns)
mappend = withAppend @ns @ns (wiring (dimsOf @ns) (wires @ns P.++ wires @ns) (wires @ns))
-- | Copying each wire, and discarding.
instance (Scalar e, KnownListOf KnownNat ns) => Comonoid (TN ns :: TNET e) where
counit = wiring (dimsOf @ns) (wires @ns) []
comult = withAppend @ns @ns (wiring (dimsOf @ns) (wires @ns) (wires @ns P.++ wires @ns))
-- | A node for each wire of a list.
wires :: forall ns. (KnownListOf KnownNat ns) => [Node]
wires = [0 .. P.length (dimsOf @ns) - 1]
instance (Scalar e, KnownListOf KnownNat ns) => CommutativeMonoid (TN ns :: TNET e)
instance (Scalar e, KnownListOf KnownNat ns) => CocommutativeComonoid (TN ns :: TNET e)
instance (Scalar e, KnownListOf KnownNat ns) => Frobenius (TN ns :: TNET e)
instance (Scalar e) => Hypergraph (TNET e)
instance (Scalar e) => CopyDiscard (TNET e)
-- | The size of an object is the product of its dimensions.
instance (Scalar e) => Sized (TNET e) where
sizeOf @(TN ns) = size @ns
-- | Two objects are isomorphic when they have the same dimensions.
instance (Scalar e) => DecidableIso (TNET e) where
isoOf @_ @(TN as) @(TN bs) =
isoFromEquality
( P.fmap
(\Refl -> Refl)
(eqListOf @KnownNat (\ @x @y -> sameNat (Proxy @x) (Proxy @y)) (listOf @KnownNat @as) (listOf @KnownNat @bs))
)