{-# LANGUAGE DerivingStrategies #-}
module Circuit.LLM.SSM
(
Aff (..),
affComp,
seqSSM,
assocScan,
assocSSM,
chunkedScan,
AffVec (..),
affCompVec,
seqSSMVec,
assocScanVec,
assocSSMVec,
runSystem,
ssmSystemVec,
multiHeadSSMSystem,
runMultiHeadSSMSystem,
runSharedInputMultiHeadSSMSystem,
ssmProcess,
ssmSystem,
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)
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)
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)
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
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
[] -> []
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)
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
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
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
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)
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)
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
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)
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
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)
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', ()))
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
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 []
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]
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
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