{-# LANGUAGE DerivingStrategies #-}

-- | Linear state-space model as a 'Process' and the associative-scan law.
--
-- A scalar linear SSM is the recurrence @h_t = a_t h_{t-1} + b_t@.  Each step
-- is an affine function @(a_t, b_t)@; composition of affine functions is
-- associative, so the whole sequence can be reduced via an associative scan
-- instead of a sequential fold.  This is the parallelisation principle behind
-- S4 / Mamba / linear attention.
module Circuit.LLM.SSM
  ( -- * Scalar affine SSM
    Aff (..),
    affComp,
    seqSSM,
    assocScan,
    assocSSM,
    chunkedScan,

    -- * Vector (harpie) affine SSM
    AffVec (..),
    affCompVec,
    seqSSMVec,
    assocScanVec,
    assocSSMVec,

    -- * System view
    runSystem,
    ssmSystemVec,

    -- * Multi-head System view (Dirichlet tensor)
    multiHeadSSMSystem,
    runMultiHeadSSMSystem,
    runSharedInputMultiHeadSSMSystem,

    -- * Process / System view
    ssmProcess,
    ssmSystem,

    -- * Centrality pair (multi-head coupling)
    coupledMultiHeadSSMSystem,
  )
where

import Circuit.Body (Body (..))
import Circuit.Poly (Mono, Poly (Tensor))
import Circuit.Process (Process (..), systemToProcess)
import Circuit.System (System, SystemT (..), monoDir, monoIn, mooreSystem, system)
import Data.List (foldl1', scanl')
import Data.Void (absurd)
import Harpie.Array (Array, zipWith)
import Prelude hiding (zipWith)

-- | An affine function @h -> a * h + b@.
data Aff = Aff
  { Aff -> Double
affA :: Double,
    Aff -> Double
affB :: Double
  }
  deriving stock (Aff -> Aff -> Bool
(Aff -> Aff -> Bool) -> (Aff -> Aff -> Bool) -> Eq Aff
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: Aff -> Aff -> Bool
== :: Aff -> Aff -> Bool
$c/= :: Aff -> Aff -> Bool
/= :: Aff -> Aff -> Bool
Eq, Int -> Aff -> ShowS
[Aff] -> ShowS
Aff -> String
(Int -> Aff -> ShowS)
-> (Aff -> String) -> ([Aff] -> ShowS) -> Show Aff
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> Aff -> ShowS
showsPrec :: Int -> Aff -> ShowS
$cshow :: Aff -> String
show :: Aff -> String
$cshowList :: [Aff] -> ShowS
showList :: [Aff] -> ShowS
Show)

-- | Composition of affine functions: @(f2 . f1)(h) = f2(f1(h))@.
--
-- >>> let f1 = Aff 2 3; f2 = Aff 5 7 in affComp f2 f1
-- Aff {affA = 10.0, affB = 22.0}
--
-- Because @(f2 . f1)(h) = 5*(2*h+3)+7 = 10*h + 22@.
affComp :: Aff -> Aff -> Aff
affComp :: Aff -> Aff -> Aff
affComp (Aff Double
a2 Double
b2) (Aff Double
a1 Double
b1) = Double -> Double -> Aff
Aff (Double
a2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
a1) (Double
a2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
b1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b2)

-- | Sequential scan of a scalar SSM from an initial state.
--
-- >>> seqSSM 0 [Aff 1 1, Aff 1 2, Aff 1 3]
-- [1.0,3.0,6.0]
seqSSM :: Double -> [Aff] -> [Double]
seqSSM :: Double -> [Aff] -> [Double]
seqSSM Double
h0 [Aff]
affs = case (Double -> Aff -> Double) -> Double -> [Aff] -> [Double]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' Double -> Aff -> Double
step Double
h0 [Aff]
affs of
  (Double
_ : [Double]
outs) -> [Double]
outs
  [] -> []
  where
    step :: Double -> Aff -> Double
step Double
h (Aff Double
a Double
b) = Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b

-- | Associative scan on affine coefficients.  Returns the list of composed
-- affine functions @[f1, f2 . f1, f3 . f2 . f1, ...]@.
--
-- The operator is reversed composition because 'scanl'' threads the accumulator
-- on the left, but we want the rightmost function applied first.
assocScan :: [Aff] -> [Aff]
assocScan :: [Aff] -> [Aff]
assocScan [Aff]
affs = case (Aff -> Aff -> Aff) -> Aff -> [Aff] -> [Aff]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' ((Aff -> Aff -> Aff) -> Aff -> Aff -> Aff
forall a b c. (a -> b -> c) -> b -> a -> c
flip Aff -> Aff -> Aff
affComp) (Double -> Double -> Aff
Aff Double
1 Double
0) [Aff]
affs of
  (Aff
_ : [Aff]
comps) -> [Aff]
comps
  [] -> []

-- | Tree-shaped associative scan with a given chunk size.
--
-- Each chunk is reduced to a summary affine function, the summaries are
-- prefix-composed, and the result is expanded back to a per-step prefix.  If
-- 'affComp' is associative, this must agree with 'assocScan' for every chunk
-- size.  This is the genuine parallelisation principle behind S4 / Mamba /
-- linear attention.
chunkedScan :: Int -> [Aff] -> [Aff]
chunkedScan :: Int -> [Aff] -> [Aff]
chunkedScan Int
_ [] = []
chunkedScan Int
k [Aff]
affs
  | Int
k Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
<= Int
1 = [Aff] -> [Aff]
assocScan [Aff]
affs
  | Bool
otherwise =
      let chunks :: [[Aff]]
chunks = Int -> [Aff] -> [[Aff]]
forall {a}. Int -> [a] -> [[a]]
chunksOf Int
k [Aff]
affs
          summaries :: [Aff]
summaries = ([Aff] -> Aff) -> [[Aff]] -> [Aff]
forall a b. (a -> b) -> [a] -> [b]
map ((Aff -> Aff -> Aff) -> [Aff] -> Aff
forall a. HasCallStack => (a -> a -> a) -> [a] -> a
foldl1' ((Aff -> Aff -> Aff) -> Aff -> Aff -> Aff
forall a b c. (a -> b -> c) -> b -> a -> c
flip Aff -> Aff -> Aff
affComp)) [[Aff]]
chunks
          summaryPrefixes :: [Aff]
summaryPrefixes = [Aff] -> [Aff]
assocScan [Aff]
summaries
          expand :: [Aff] -> Aff -> [Aff]
expand [Aff]
chunk Aff
prevPref =
            let inner :: [Aff]
inner = case (Aff -> Aff -> Aff) -> Aff -> [Aff] -> [Aff]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' ((Aff -> Aff -> Aff) -> Aff -> Aff -> Aff
forall a b c. (a -> b -> c) -> b -> a -> c
flip Aff -> Aff -> Aff
affComp) (Double -> Double -> Aff
Aff Double
1 Double
0) [Aff]
chunk of (Aff
_ : [Aff]
xs) -> [Aff]
xs; [] -> []
             in (Aff -> Aff) -> [Aff] -> [Aff]
forall a b. (a -> b) -> [a] -> [b]
map (\Aff
x -> Aff -> Aff -> Aff
affComp Aff
x Aff
prevPref) [Aff]
inner
       in [[Aff]] -> [Aff]
forall (t :: * -> *) a. Foldable t => t [a] -> [a]
concat [[Aff] -> Aff -> [Aff]
expand [Aff]
chunk Aff
prevPref | ([Aff]
chunk, Aff
prevPref) <- [[Aff]] -> [Aff] -> [([Aff], Aff)]
forall a b. [a] -> [b] -> [(a, b)]
zip [[Aff]]
chunks (Double -> Double -> Aff
Aff Double
1 Double
0 Aff -> [Aff] -> [Aff]
forall a. a -> [a] -> [a]
: [Aff]
summaryPrefixes)]
  where
    chunksOf :: Int -> [a] -> [[a]]
chunksOf Int
_ [] = []
    chunksOf Int
n [a]
xs = Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
take Int
n [a]
xs [a] -> [[a]] -> [[a]]
forall a. a -> [a] -> [a]
: Int -> [a] -> [[a]]
chunksOf Int
n (Int -> [a] -> [a]
forall a. Int -> [a] -> [a]
drop Int
n [a]
xs)

-- | Apply the associatively-scanned affine functions to an initial state.
assocSSM :: Double -> [Aff] -> [Double]
assocSSM :: Double -> [Aff] -> [Double]
assocSSM Double
h0 = (Aff -> Double) -> [Aff] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (\(Aff Double
a Double
b) -> Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h0 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b) ([Aff] -> [Double]) -> ([Aff] -> [Aff]) -> [Aff] -> [Double]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. [Aff] -> [Aff]
assocScan

-- | A 'System' whose state is the hidden state @h@ and whose output is @h@.
-- Input is the affine coefficient pair @(a_t, b_t)@; the initial state @h0@ is
-- supplied when converting to a 'Process' or running directly.
ssmSystem :: System (->) Double (Mono Aff Double)
ssmSystem :: System (->) Double (Mono Aff Double)
ssmSystem = (Double -> Aff -> Double)
-> (Double -> Double) -> System (->) Double (Mono Aff Double)
forall s a b. (s -> a -> s) -> (s -> b) -> System (->) s (Mono a b)
mooreSystem Double -> Aff -> Double
step Double -> Double
forall {p}. p -> p
extract
  where
    step :: Double -> Aff -> Double
step Double
h (Aff Double
a Double
b) = Double
a Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b
    extract :: p -> p
extract p
h = p
h

-- | A 'Process' whose state is the hidden state @h@ and whose output is @h@.
-- Input is the affine coefficient pair @(a_t, b_t)@.
--
-- This is the first-input-seeded presentation with @h0 = 0@.  Use
-- 'ssmSystem' with 'systemToProcess' when you need a non-zero seed.
ssmProcess :: Process Aff Double
ssmProcess :: Process Aff Double
ssmProcess = Double
-> (Double -> Double)
-> System (->) Double (Mono Aff Double)
-> Process Aff Double
forall s b a.
s -> (s -> b) -> System (->) s (Mono a b) -> Process a b
systemToProcess Double
0 Double -> Double
forall {p}. p -> p
id System (->) Double (Mono Aff Double)
ssmSystem

-- ---------------------------------------------------------------------------
-- Vector (harpie) affine SSM — diagonal-matrix state
-- ---------------------------------------------------------------------------

-- | Elementwise affine function on a harpie 1-D array:
-- @h_i -> a_i * h_i + b_i@.  This is the diagonal-matrix special case of the
-- full matrix SSM; it already exercises the associative-scan law on arrays.
data AffVec = AffVec
  { AffVec -> Array Double
affAVec :: Array Double,
    AffVec -> Array Double
affBVec :: Array Double
  }
  deriving stock (AffVec -> AffVec -> Bool
(AffVec -> AffVec -> Bool)
-> (AffVec -> AffVec -> Bool) -> Eq AffVec
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: AffVec -> AffVec -> Bool
== :: AffVec -> AffVec -> Bool
$c/= :: AffVec -> AffVec -> Bool
/= :: AffVec -> AffVec -> Bool
Eq, Int -> AffVec -> ShowS
[AffVec] -> ShowS
AffVec -> String
(Int -> AffVec -> ShowS)
-> (AffVec -> String) -> ([AffVec] -> ShowS) -> Show AffVec
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> AffVec -> ShowS
showsPrec :: Int -> AffVec -> ShowS
$cshow :: AffVec -> String
show :: AffVec -> String
$cshowList :: [AffVec] -> ShowS
showList :: [AffVec] -> ShowS
Show)

-- | Elementwise composition of diagonal affine functions.
affCompVec :: AffVec -> AffVec -> AffVec
affCompVec :: AffVec -> AffVec -> AffVec
affCompVec (AffVec Array Double
a2 Array Double
b2) (AffVec Array Double
a1 Array Double
b1) =
  Array Double -> Array Double -> AffVec
AffVec ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a2 Array Double
a1) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a2 Array Double
b1) Array Double
b2)

-- | Sequential scan of a vector SSM.
seqSSMVec :: Array Double -> [AffVec] -> [Array Double]
seqSSMVec :: Array Double -> [AffVec] -> [Array Double]
seqSSMVec Array Double
h0 [AffVec]
affs = case (Array Double -> AffVec -> Array Double)
-> Array Double -> [AffVec] -> [Array Double]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' Array Double -> AffVec -> Array Double
step Array Double
h0 [AffVec]
affs of
  (Array Double
_ : [Array Double]
outs) -> [Array Double]
outs
  [] -> []
  where
    step :: Array Double -> AffVec -> Array Double
step Array Double
h (AffVec Array Double
a Array Double
b) = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a Array Double
h) Array Double
b

-- | Associative scan on vector affine coefficients.  The shape is inherited
-- from the first input.
assocScanVec :: [AffVec] -> [AffVec]
assocScanVec :: [AffVec] -> [AffVec]
assocScanVec [] = []
assocScanVec (AffVec
firstStep : [AffVec]
rest) =
  case (AffVec -> AffVec -> AffVec) -> AffVec -> [AffVec] -> [AffVec]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' ((AffVec -> AffVec -> AffVec) -> AffVec -> AffVec -> AffVec
forall a b c. (a -> b -> c) -> b -> a -> c
flip AffVec -> AffVec -> AffVec
affCompVec) (Array Double -> Array Double -> AffVec
AffVec Array Double
ones Array Double
zeros) (AffVec
firstStep AffVec -> [AffVec] -> [AffVec]
forall a. a -> [a] -> [a]
: [AffVec]
rest) of
    (AffVec
_ : [AffVec]
comps) -> [AffVec]
comps
    [] -> []
  where
    ones :: Array Double
ones = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith (\Double
_ Double
_ -> Double
1) (AffVec -> Array Double
affAVec AffVec
firstStep) (AffVec -> Array Double
affAVec AffVec
firstStep)
    zeros :: Array Double
zeros = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith (\Double
_ Double
_ -> Double
0) (AffVec -> Array Double
affBVec AffVec
firstStep) (AffVec -> Array Double
affBVec AffVec
firstStep)

-- | Apply the associatively-scanned vector affine functions to an initial state.
assocSSMVec :: Array Double -> [AffVec] -> [Array Double]
assocSSMVec :: Array Double -> [AffVec] -> [Array Double]
assocSSMVec Array Double
h0 = (AffVec -> Array Double) -> [AffVec] -> [Array Double]
forall a b. (a -> b) -> [a] -> [b]
map (\(AffVec Array Double
a Array Double
b) -> (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a Array Double
h0) Array Double
b) ([AffVec] -> [Array Double])
-> ([AffVec] -> [AffVec]) -> [AffVec] -> [Array Double]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. [AffVec] -> [AffVec]
assocScanVec

-- ---------------------------------------------------------------------------
-- System view: linear SSM as a Moore machine
-- ---------------------------------------------------------------------------

-- | Run a deterministic 'System' with a monomial interface over a list of
-- inputs.  This is the same semantics as 'Circuit.Process.scan', but stated
-- directly on 'System'.
runSystem :: System (->) s (Mono i o) -> s -> [i] -> ([o], s)
runSystem :: forall s i o. System (->) s (Mono i o) -> s -> [i] -> ([o], s)
runSystem (SystemT (Body (s, Dir (Mono i o)) -> (s, Pos (Mono i o))
sys)) s
s0 [i]
is = s -> [i] -> [o] -> ([o], s)
go s
s0 [i]
is []
  where
    go :: s -> [i] -> [o] -> ([o], s)
go s
s [] [o]
acc = ([o] -> [o]
forall a. [a] -> [a]
reverse [o]
acc, s
s)
    go s
s (i
i : [i]
iss) [o]
acc =
      let (s
s', (o
o, ())) = (s, Dir (Mono i o)) -> (s, Pos (Mono i o))
sys (s
s, i -> Dir (Mono i (ZonkAny 3))
forall i o. i -> Dir (Mono i o)
monoIn i
i)
       in s -> [i] -> [o] -> ([o], s)
go s
s' [i]
iss (o
o o -> [o] -> [o]
forall a. a -> [a] -> [a]
: [o]
acc)

-- | Vector SSM as a 'System (->)' with harpie state, input 'AffVec', and full
-- state observation.
ssmSystemVec :: System (->) (Array Double) (Mono AffVec (Array Double))
ssmSystemVec :: System (->) (Array Double) (Mono AffVec (Array Double))
ssmSystemVec = ((Array Double, Dir (Mono AffVec (Array Double)))
 -> (Array Double, Pos (Mono AffVec (Array Double))))
-> System (->) (Array Double) (Mono AffVec (Array Double))
forall (arr :: * -> * -> *) s (p :: Poly).
arr (s, Dir p) (s, Pos p) -> System arr s p
system (((Array Double, Dir (Mono AffVec (Array Double)))
  -> (Array Double, Pos (Mono AffVec (Array Double))))
 -> System (->) (Array Double) (Mono AffVec (Array Double)))
-> ((Array Double, Dir (Mono AffVec (Array Double)))
    -> (Array Double, Pos (Mono AffVec (Array Double))))
-> System (->) (Array Double) (Mono AffVec (Array Double))
forall a b. (a -> b) -> a -> b
$ \(Array Double
h, Dir (Mono AffVec (Array Double))
d) ->
  let AffVec Array Double
a Array Double
b = Dir (Mono AffVec (ZonkAny 2)) -> AffVec
forall i o. Dir (Mono i o) -> i
monoDir Dir (Mono AffVec (ZonkAny 2))
Dir (Mono AffVec (Array Double))
d
      h' :: Array Double
h' = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a Array Double
h) Array Double
b
   in (Array Double
h', (Array Double
h', ()))

-- ---------------------------------------------------------------------------
-- Multi-head SSM as a Tensor-polynomial System
-- ---------------------------------------------------------------------------

-- | Two independent vector SSM heads packaged as a single 'System' over the
-- Dirichlet tensor @Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double))@.
--
-- The two heads share the same input /direction type/ ('AffVec') but receive
-- independent direction values.  This is the right polynomial for parallel
-- layers: both heads fire on the same tick, each with its own input.  A
-- cartesian 'Prod' would force a choice between heads via @Either@ directions.
multiHeadSSMSystem ::
  System
    (->)
    (Array Double, Array Double)
    (Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
multiHeadSSMSystem :: System
  (->)
  (Array Double, Array Double)
  ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
multiHeadSSMSystem = (((Array Double, Array Double),
  Dir
    ('Tensor
       (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
 -> ((Array Double, Array Double),
     Pos
       ('Tensor
          (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))))
-> System
     (->)
     (Array Double, Array Double)
     ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
forall (arr :: * -> * -> *) s (p :: Poly).
arr (s, Dir p) (s, Pos p) -> System arr s p
system ((((Array Double, Array Double),
   Dir
     ('Tensor
        (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
  -> ((Array Double, Array Double),
      Pos
        ('Tensor
           (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))))
 -> System
      (->)
      (Array Double, Array Double)
      ('Tensor
         (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
-> (((Array Double, Array Double),
     Dir
       ('Tensor
          (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
    -> ((Array Double, Array Double),
        Pos
          ('Tensor
             (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))))
-> System
     (->)
     (Array Double, Array Double)
     ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
forall a b. (a -> b) -> a -> b
$ \case
  ((Array Double
h1, Array Double
h2), (Right AffVec
aff1, Right AffVec
aff2)) ->
    let AffVec Array Double
a1 Array Double
b1 = AffVec
aff1
        AffVec Array Double
a2 Array Double
b2 = AffVec
aff2
        h1' :: Array Double
h1' = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a1 Array Double
h1) Array Double
b1
        h2' :: Array Double
h2' = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a2 Array Double
h2) Array Double
b2
     in ((Array Double
h1', Array Double
h2'), ((Array Double
h1', ()), (Array Double
h2', ())))
  ((Array Double, Array Double)
_, (Left Void
v, Either Void AffVec
_)) -> Void
-> ((Array Double, Array Double),
    ((Array Double, ()), (Array Double, ())))
forall a. Void -> a
absurd Void
v
  ((Array Double, Array Double)
_, (Either Void AffVec
_, Left Void
v)) -> Void
-> ((Array Double, Array Double),
    ((Array Double, ()), (Array Double, ())))
forall a. Void -> a
absurd Void
v

-- | Run the multi-head SSM with independent per-head inputs.
--
-- Returns the pair of head outputs at each step and the final pair of states.
runMultiHeadSSMSystem ::
  (Array Double, Array Double) ->
  [(AffVec, AffVec)] ->
  ([(Array Double, Array Double)], (Array Double, Array Double))
runMultiHeadSSMSystem :: (Array Double, Array Double)
-> [(AffVec, AffVec)]
-> ([(Array Double, Array Double)], (Array Double, Array Double))
runMultiHeadSSMSystem (Array Double, Array Double)
s0 [(AffVec, AffVec)]
affPairs =
  let SystemT (Body ((Array Double, Array Double),
 Dir
   ('Tensor
      (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
-> ((Array Double, Array Double),
    Pos
      ('Tensor
         (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
f) = System
  (->)
  (Array Double, Array Double)
  ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
multiHeadSSMSystem
      go :: (Array Double, Array Double)
-> [(AffVec, AffVec)]
-> [(Array Double, Array Double)]
-> ([(Array Double, Array Double)], (Array Double, Array Double))
go (Array Double, Array Double)
s [] [(Array Double, Array Double)]
acc = ([(Array Double, Array Double)] -> [(Array Double, Array Double)]
forall a. [a] -> [a]
reverse [(Array Double, Array Double)]
acc, (Array Double, Array Double)
s)
      go (Array Double
h1, Array Double
h2) ((AffVec
aff1, AffVec
aff2) : [(AffVec, AffVec)]
affs') [(Array Double, Array Double)]
acc =
        let ((Array Double
h1', Array Double
h2'), ((Array Double
o1, ()), (Array Double
o2, ()))) = ((Array Double, Array Double),
 (Either Void AffVec, Either Void AffVec))
-> ((Array Double, Array Double),
    ((Array Double, ()), (Array Double, ())))
f ((Array Double
h1, Array Double
h2), (AffVec -> Dir (Mono AffVec (ZonkAny 0))
forall i o. i -> Dir (Mono i o)
monoIn AffVec
aff1, AffVec -> Dir (Mono AffVec (ZonkAny 1))
forall i o. i -> Dir (Mono i o)
monoIn AffVec
aff2))
         in (Array Double, Array Double)
-> [(AffVec, AffVec)]
-> [(Array Double, Array Double)]
-> ([(Array Double, Array Double)], (Array Double, Array Double))
go (Array Double
h1', Array Double
h2') [(AffVec, AffVec)]
affs' ((Array Double
o1, Array Double
o2) (Array Double, Array Double)
-> [(Array Double, Array Double)] -> [(Array Double, Array Double)]
forall a. a -> [a] -> [a]
: [(Array Double, Array Double)]
acc)
   in (Array Double, Array Double)
-> [(AffVec, AffVec)]
-> [(Array Double, Array Double)]
-> ([(Array Double, Array Double)], (Array Double, Array Double))
go (Array Double, Array Double)
s0 [(AffVec, AffVec)]
affPairs []

-- | Run the multi-head SSM with the /same/ input supplied to both heads.
--
-- This is the diagonal shared-input case; building it from independent inputs
-- requires an explicit copy of the input direction.
runSharedInputMultiHeadSSMSystem ::
  (Array Double, Array Double) ->
  [AffVec] ->
  ([(Array Double, Array Double)], (Array Double, Array Double))
runSharedInputMultiHeadSSMSystem :: (Array Double, Array Double)
-> [AffVec]
-> ([(Array Double, Array Double)], (Array Double, Array Double))
runSharedInputMultiHeadSSMSystem (Array Double, Array Double)
s0 [AffVec]
affs = (Array Double, Array Double)
-> [(AffVec, AffVec)]
-> ([(Array Double, Array Double)], (Array Double, Array Double))
runMultiHeadSSMSystem (Array Double, Array Double)
s0 [(AffVec
aff, AffVec
aff) | AffVec
aff <- [AffVec]
affs]

-- | Two-headed SSM with a cross-head coupling: head 1 reads head 2's state.
--
-- This breaks premonoidal centrality: the order in which the two heads are
-- threaded through the shared medium matters, because head 1's update depends
-- on head 2's state.  It is the "flip" oracle for the multi-head centrality
-- pair.
coupledMultiHeadSSMSystem ::
  System
    (->)
    (Array Double, Array Double)
    (Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
coupledMultiHeadSSMSystem :: System
  (->)
  (Array Double, Array Double)
  ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
coupledMultiHeadSSMSystem = (((Array Double, Array Double),
  Dir
    ('Tensor
       (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
 -> ((Array Double, Array Double),
     Pos
       ('Tensor
          (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))))
-> System
     (->)
     (Array Double, Array Double)
     ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
forall (arr :: * -> * -> *) s (p :: Poly).
arr (s, Dir p) (s, Pos p) -> System arr s p
system ((((Array Double, Array Double),
   Dir
     ('Tensor
        (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
  -> ((Array Double, Array Double),
      Pos
        ('Tensor
           (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))))
 -> System
      (->)
      (Array Double, Array Double)
      ('Tensor
         (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
-> (((Array Double, Array Double),
     Dir
       ('Tensor
          (Mono AffVec (Array Double)) (Mono AffVec (Array Double))))
    -> ((Array Double, Array Double),
        Pos
          ('Tensor
             (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))))
-> System
     (->)
     (Array Double, Array Double)
     ('Tensor (Mono AffVec (Array Double)) (Mono AffVec (Array Double)))
forall a b. (a -> b) -> a -> b
$ \case
  ((Array Double
h1, Array Double
h2), (Right AffVec
aff1, Right AffVec
aff2)) ->
    let AffVec Array Double
a1 Array Double
b1 = AffVec
aff1
        AffVec Array Double
a2 Array Double
b2 = AffVec
aff2
        -- Head 1 receives a cross-term from head 2's state.
        cross :: Array Double
cross = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
h2 ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith (\Double
_ Double
_ -> Double
0.1) Array Double
h2 Array Double
h2)
        h1' :: Array Double
h1' = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a1 Array Double
h1) Array Double
b1) Array Double
cross
        h2' :: Array Double
h2' = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) ((Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
a2 Array Double
h2) Array Double
b2
     in ((Array Double
h1', Array Double
h2'), ((Array Double
h1', ()), (Array Double
h2', ())))
  ((Array Double, Array Double)
_, (Left Void
v, Either Void AffVec
_)) -> Void
-> ((Array Double, Array Double),
    ((Array Double, ()), (Array Double, ())))
forall a. Void -> a
absurd Void
v
  ((Array Double, Array Double)
_, (Either Void AffVec
_, Left Void
v)) -> Void
-> ((Array Double, Array Double),
    ((Array Double, ()), (Array Double, ())))
forall a. Void -> a
absurd Void
v