{-# LANGUAGE PatternSynonyms #-}
{-# LANGUAGE RebindableSyntax #-}
{-# OPTIONS_GHC -Wno-pattern-namespace-specifier #-}
module Circuit.Stats.Diff
(
GradInputs (..),
variable,
variables,
DiffProcess (..),
StdState (..),
diffFold,
diffScan,
DiffSystem (..),
diffSystemAsDiffProcess,
diffProcessAsDiffSystem,
diffSystemScan,
diffSystemFold,
gradFold,
gradScan,
constant,
netFold,
netScan,
onlineDiff,
maDiff,
sqmaDiff,
stdDiff,
onlineDiffProcess,
maDiffProcess,
sqmaDiffProcess,
stdDiffProcess,
delay1Diff,
diffDiff,
)
where
import Circuit.Diff (Diff, Diff', runDiff, pattern Diff)
import Circuit.Stats (Averager (..), Process, fold, ma, online, scan, sqma, std, pattern A)
import Data.List (length)
import Data.Vector.Unboxed qualified as VU
import Harpie.Array (Array)
import Harpie.Array qualified as HA
import NumHask.Prelude hiding (fold, length)
newtype GradInputs a = GradInputs
{ forall a. GradInputs a -> Array a
unGradInputs :: Array a
}
deriving (GradInputs a -> GradInputs a -> Bool
(GradInputs a -> GradInputs a -> Bool)
-> (GradInputs a -> GradInputs a -> Bool) -> Eq (GradInputs a)
forall a. Eq a => GradInputs a -> GradInputs a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => GradInputs a -> GradInputs a -> Bool
== :: GradInputs a -> GradInputs a -> Bool
$c/= :: forall a. Eq a => GradInputs a -> GradInputs a -> Bool
/= :: GradInputs a -> GradInputs a -> Bool
Eq, Int -> GradInputs a -> ShowS
[GradInputs a] -> ShowS
GradInputs a -> String
(Int -> GradInputs a -> ShowS)
-> (GradInputs a -> String)
-> ([GradInputs a] -> ShowS)
-> Show (GradInputs a)
forall a. Show a => Int -> GradInputs a -> ShowS
forall a. Show a => [GradInputs a] -> ShowS
forall a. Show a => GradInputs a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> GradInputs a -> ShowS
showsPrec :: Int -> GradInputs a -> ShowS
$cshow :: forall a. Show a => GradInputs a -> String
show :: GradInputs a -> String
$cshowList :: forall a. Show a => [GradInputs a] -> ShowS
showList :: [GradInputs a] -> ShowS
Show)
broadcastShape :: Array a -> Array a -> [Int]
broadcastShape :: forall a. Array a -> Array a -> [Int]
broadcastShape Array a
xs Array a
ys = [Int] -> [Int] -> [Int]
forall {a}. Ord a => [a] -> [a] -> [a]
go (Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
HA.shape Array a
xs)) (Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
HA.shape Array a
ys))
where
go :: [a] -> [a] -> [a]
go [] [a]
bs = [a]
bs
go [a]
as [] = [a]
as
go (a
a : [a]
as) (a
b : [a]
bs) = a -> a -> a
forall a. Ord a => a -> a -> a
max a
a a
b a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a] -> [a] -> [a]
go [a]
as [a]
bs
instance (Additive a) => Additive (GradInputs a) where
zero :: GradInputs a
zero = Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ([Int] -> a -> Array a
forall a. [Int] -> a -> Array a
HA.konst [Int
0] a
forall a. Additive a => a
zero)
GradInputs Array a
xs + :: GradInputs a -> GradInputs a -> GradInputs a
+ GradInputs Array a
ys
| Array a -> Vector Int
forall a. Array a -> Vector Int
HA.shape Array a
xs Vector Int -> Vector Int -> Bool
forall a. Eq a => a -> a -> Bool
== Array a -> Vector Int
forall a. Array a -> Vector Int
HA.shape Array a
ys = Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ((a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
HA.zipWith a -> a -> a
forall a. Additive a => a -> a -> a
(+) Array a
xs Array a
ys)
| Bool
otherwise =
let s :: [Int]
s = Array a -> Array a -> [Int]
forall a. Array a -> Array a -> [Int]
broadcastShape Array a
xs Array a
ys
in Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ((a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
HA.zipWith a -> a -> a
forall a. Additive a => a -> a -> a
(+) (a -> [Int] -> Array a -> Array a
forall a. a -> [Int] -> Array a -> Array a
HA.pad a
forall a. Additive a => a
zero [Int]
s Array a
xs) (a -> [Int] -> Array a -> Array a
forall a. a -> [Int] -> Array a -> Array a
HA.pad a
forall a. Additive a => a
zero [Int]
s Array a
ys))
instance (Subtractive a) => Subtractive (GradInputs a) where
negate :: GradInputs a -> GradInputs a
negate (GradInputs Array a
xs) = Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ((a -> a) -> Array a -> Array a
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap a -> a
forall a. Subtractive a => a -> a
negate Array a
xs)
GradInputs Array a
xs - :: GradInputs a -> GradInputs a -> GradInputs a
- GradInputs Array a
ys
| Array a -> Vector Int
forall a. Array a -> Vector Int
HA.shape Array a
xs Vector Int -> Vector Int -> Bool
forall a. Eq a => a -> a -> Bool
== Array a -> Vector Int
forall a. Array a -> Vector Int
HA.shape Array a
ys = Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ((a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
HA.zipWith (-) Array a
xs Array a
ys)
| Bool
otherwise =
let s :: [Int]
s = Array a -> Array a -> [Int]
forall a. Array a -> Array a -> [Int]
broadcastShape Array a
xs Array a
ys
in Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ((a -> a -> a) -> Array a -> Array a -> Array a
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
HA.zipWith (-) (a -> [Int] -> Array a -> Array a
forall a. a -> [Int] -> Array a -> Array a
HA.pad a
forall a. Additive a => a
zero [Int]
s Array a
xs) (a -> [Int] -> Array a -> Array a
forall a. a -> [Int] -> Array a -> Array a
HA.pad a
forall a. Additive a => a
zero [Int]
s Array a
ys))
variable :: (Additive a) => Int -> Int -> Diff' (GradInputs a) a
variable :: forall a. Additive a => Int -> Int -> Diff' (GradInputs a) a
variable Int
n Int
i = (GradInputs a -> (a, a -> GradInputs a))
-> Diff () (GradInputs a) a
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((GradInputs a -> (a, a -> GradInputs a))
-> Diff () (GradInputs a) a)
-> (GradInputs a -> (a, a -> GradInputs a))
-> Diff () (GradInputs a) a
forall a b. (a -> b) -> a -> b
$ \GradInputs a
s ->
( Array a -> [Int] -> a
forall a. Array a -> [Int] -> a
HA.index (GradInputs a -> Array a
forall a. GradInputs a -> Array a
unGradInputs GradInputs a
s) [Int
i],
\a
da -> Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ([Int] -> (a -> a) -> Array a -> Array a
forall a. [Int] -> (a -> a) -> Array a -> Array a
HA.modify [Int
i] (a -> a -> a
forall a b. a -> b -> a
const a
da) ([Int] -> a -> Array a
forall a. [Int] -> a -> Array a
HA.konst [Int
n] a
forall a. Additive a => a
zero))
)
variables :: (Additive a) => [a] -> [Diff' (GradInputs a) a]
variables :: forall a. Additive a => [a] -> [Diff' (GradInputs a) a]
variables [a]
xs =
let n :: Int
n = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
xs
in [Int -> Int -> Diff' (GradInputs a) a
forall a. Additive a => Int -> Int -> Diff' (GradInputs a) a
variable Int
n Int
i | Int
i <- [Int
0 .. Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1]]
gradFold ::
(Additive a) =>
Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b) ->
[a] ->
(b, b -> [a])
gradFold :: forall a b.
Additive a =>
Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
-> [a] -> (b, b -> [a])
gradFold Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
_ [] = String -> (b, b -> [a])
forall a. HasCallStack => String -> a
error String
"gradFold: empty list"
gradFold Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
m [a]
xs =
case Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
-> [Diff' (GradInputs a) a] -> Maybe (Diff' (GradInputs a) b)
forall a b. Process a b -> [a] -> Maybe b
fold Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
m ([a] -> [Diff' (GradInputs a) a]
forall a. Additive a => [a] -> [Diff' (GradInputs a) a]
variables [a]
xs) of
Maybe (Diff' (GradInputs a) b)
Nothing -> String -> (b, b -> [a])
forall a. HasCallStack => String -> a
error String
"gradFold: empty fold"
Just Diff' (GradInputs a) b
ydiff ->
let (b
y, b -> GradInputs a
pb) = Diff' (GradInputs a) b -> GradInputs a -> (b, b -> GradInputs a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' (GradInputs a) b
ydiff (Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ([Int] -> [a] -> Array a
forall t a. FromVector t a => [Int] -> t -> Array a
HA.array [[a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
xs] [a]
xs))
in (b
y, Array a -> [a]
forall t a. FromArray t a => Array a -> t
HA.arrayAs (Array a -> [a]) -> (b -> Array a) -> b -> [a]
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. GradInputs a -> Array a
forall a. GradInputs a -> Array a
unGradInputs (GradInputs a -> Array a) -> (b -> GradInputs a) -> b -> Array a
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. b -> GradInputs a
pb)
gradScan ::
(Additive a) =>
Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b) ->
[a] ->
([b], [b] -> [a])
gradScan :: forall a b.
Additive a =>
Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
-> [a] -> ([b], [b] -> [a])
gradScan Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
_ [] = ([], [a] -> [b] -> [a]
forall a b. a -> b -> a
const [])
gradScan Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
m [a]
xs =
let ys :: [Diff' (GradInputs a) b]
ys = Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
-> [Diff' (GradInputs a) a] -> [Diff' (GradInputs a) b]
forall a b. Process a b -> [a] -> [b]
scan Process (Diff' (GradInputs a) a) (Diff' (GradInputs a) b)
m ([a] -> [Diff' (GradInputs a) a]
forall a. Additive a => [a] -> [Diff' (GradInputs a) a]
variables [a]
xs)
n :: Int
n = [a] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [a]
xs
pbs :: [(b, b -> GradInputs a)]
pbs = (Diff' (GradInputs a) b -> (b, b -> GradInputs a))
-> [Diff' (GradInputs a) b] -> [(b, b -> GradInputs a)]
forall a b. (a -> b) -> [a] -> [b]
map (Diff' (GradInputs a) b -> GradInputs a -> (b, b -> GradInputs a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
`runDiff` Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ([Int] -> [a] -> Array a
forall t a. FromVector t a => [Int] -> t -> Array a
HA.array [Int
n] [a]
xs)) [Diff' (GradInputs a) b]
ys
z :: GradInputs a
z = Array a -> GradInputs a
forall a. Array a -> GradInputs a
GradInputs ([Int] -> a -> Array a
forall a. [Int] -> a -> Array a
HA.konst [Int
n] a
forall a. Additive a => a
zero)
pullback :: [b] -> [a]
pullback [b]
dbs =
Array a -> [a]
forall t a. FromArray t a => Array a -> t
HA.arrayAs (Array a -> [a])
-> (GradInputs a -> Array a) -> GradInputs a -> [a]
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. GradInputs a -> Array a
forall a. GradInputs a -> Array a
unGradInputs (GradInputs a -> [a]) -> GradInputs a -> [a]
forall a b. (a -> b) -> a -> b
$
(GradInputs a -> GradInputs a -> GradInputs a)
-> GradInputs a -> [GradInputs a] -> GradInputs a
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' GradInputs a -> GradInputs a -> GradInputs a
forall a. Additive a => a -> a -> a
(+) GradInputs a
z [b -> GradInputs a
pb b
db | (b -> GradInputs a
pb, b
db) <- [b -> GradInputs a] -> [b] -> [(b -> GradInputs a, b)]
forall a b. [a] -> [b] -> [(a, b)]
zip ((b, b -> GradInputs a) -> b -> GradInputs a
forall a b. (a, b) -> b
snd ((b, b -> GradInputs a) -> b -> GradInputs a)
-> [(b, b -> GradInputs a)] -> [b -> GradInputs a]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [(b, b -> GradInputs a)]
pbs) [b]
dbs]
in ((b, b -> GradInputs a) -> b
forall a b. (a, b) -> a
fst ((b, b -> GradInputs a) -> b) -> [(b, b -> GradInputs a)] -> [b]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [(b, b -> GradInputs a)]
pbs, [b] -> [a]
pullback)
constant :: (Additive p) => a -> Diff tag p a
constant :: forall {k} p a (tag :: k). Additive p => a -> Diff tag p a
constant a
x = (p -> (a, a -> p)) -> Diff tag p a
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a, a -> p) -> p -> (a, a -> p)
forall a b. a -> b -> a
const (a
x, p -> a -> p
forall a b. a -> b -> a
const p
forall a. Additive a => a
zero))
netFold ::
(Additive p) =>
Process (Diff tag p a) (Diff tag p b) ->
p ->
[a] ->
(b, b -> p)
netFold :: forall {k} p (tag :: k) a b.
Additive p =>
Process (Diff tag p a) (Diff tag p b) -> p -> [a] -> (b, b -> p)
netFold Process (Diff tag p a) (Diff tag p b)
_ p
_ [] = String -> (b, b -> p)
forall a. HasCallStack => String -> a
error String
"netFold: empty list"
netFold Process (Diff tag p a) (Diff tag p b)
m p
p [a]
xs =
case Process (Diff tag p a) (Diff tag p b)
-> [Diff tag p a] -> Maybe (Diff tag p b)
forall a b. Process a b -> [a] -> Maybe b
fold Process (Diff tag p a) (Diff tag p b)
m ((a -> Diff tag p a) -> [a] -> [Diff tag p a]
forall a b. (a -> b) -> [a] -> [b]
map a -> Diff tag p a
forall {k} p a (tag :: k). Additive p => a -> Diff tag p a
constant [a]
xs) of
Maybe (Diff tag p b)
Nothing -> String -> (b, b -> p)
forall a. HasCallStack => String -> a
error String
"netFold: empty fold"
Just Diff tag p b
ydiff -> Diff tag p b -> p -> (b, b -> p)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff tag p b
ydiff p
p
netScan ::
(Additive p) =>
Process (Diff tag p a) (Diff tag p b) ->
p ->
[a] ->
([b], [b] -> p)
netScan :: forall {k} p (tag :: k) a b.
Additive p =>
Process (Diff tag p a) (Diff tag p b)
-> p -> [a] -> ([b], [b] -> p)
netScan Process (Diff tag p a) (Diff tag p b)
_ p
_ [] = ([], p -> [b] -> p
forall a b. a -> b -> a
const p
forall a. Additive a => a
zero)
netScan Process (Diff tag p a) (Diff tag p b)
m p
p [a]
xs =
let ys :: [Diff tag p b]
ys = Process (Diff tag p a) (Diff tag p b)
-> [Diff tag p a] -> [Diff tag p b]
forall a b. Process a b -> [a] -> [b]
scan Process (Diff tag p a) (Diff tag p b)
m ((a -> Diff tag p a) -> [a] -> [Diff tag p a]
forall a b. (a -> b) -> [a] -> [b]
map a -> Diff tag p a
forall {k} p a (tag :: k). Additive p => a -> Diff tag p a
constant [a]
xs)
pbs :: [(b, b -> p)]
pbs = (Diff tag p b -> (b, b -> p)) -> [Diff tag p b] -> [(b, b -> p)]
forall a b. (a -> b) -> [a] -> [b]
map (Diff tag p b -> p -> (b, b -> p)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
`runDiff` p
p) [Diff tag p b]
ys
pullback :: [b] -> p
pullback [b]
dbs =
(p -> p -> p) -> p -> [p] -> p
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' p -> p -> p
forall a. Additive a => a -> a -> a
(+) p
forall a. Additive a => a
zero [b -> p
pb b
db | (b -> p
pb, b
db) <- [b -> p] -> [b] -> [(b -> p, b)]
forall a b. [a] -> [b] -> [(a, b)]
zip ((b, b -> p) -> b -> p
forall a b. (a, b) -> b
snd ((b, b -> p) -> b -> p) -> [(b, b -> p)] -> [b -> p]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [(b, b -> p)]
pbs) [b]
dbs]
in ((b, b -> p) -> b
forall a b. (a, b) -> a
fst ((b, b -> p) -> b) -> [(b, b -> p)] -> [b]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> [(b, b -> p)]
pbs, [b] -> p
pullback)
onlineDiff ::
(Additive p, Subtractive b, Divisive b) =>
(Diff tag p a -> Diff tag p b) ->
(Diff tag p b -> Diff tag p b) ->
Process (Diff tag p a) (Diff tag p b)
onlineDiff :: forall {k} p b (tag :: k) a.
(Additive p, Subtractive b, Divisive b) =>
(Diff tag p a -> Diff tag p b)
-> (Diff tag p b -> Diff tag p b)
-> Process (Diff tag p a) (Diff tag p b)
onlineDiff = (Diff tag p a -> Diff tag p b)
-> (Diff tag p b -> Diff tag p b)
-> Process (Diff tag p a) (Diff tag p b)
forall b a.
(Divisive b, Additive b) =>
(a -> b) -> (b -> b) -> Process a b
online
maDiff ::
(Additive p, Subtractive b, Divisive b) =>
Diff tag p b ->
Process (Diff tag p b) (Diff tag p b)
maDiff :: forall {k} p b (tag :: k).
(Additive p, Subtractive b, Divisive b) =>
Diff tag p b -> Process (Diff tag p b) (Diff tag p b)
maDiff = Diff tag p b -> Process (Diff tag p b) (Diff tag p b)
forall a. (Divisive a, Additive a) => a -> Process a a
ma
sqmaDiff ::
(Additive p, Subtractive b, Divisive b) =>
Diff tag p b ->
Process (Diff tag p b) (Diff tag p b)
sqmaDiff :: forall {k} p b (tag :: k).
(Additive p, Subtractive b, Divisive b) =>
Diff tag p b -> Process (Diff tag p b) (Diff tag p b)
sqmaDiff = Diff tag p b -> Process (Diff tag p b) (Diff tag p b)
forall a. (Divisive a, Additive a) => a -> Process a a
sqma
stdDiff ::
(Subtractive p, ExpField b) =>
Diff tag p b ->
Process (Diff tag p b) (Diff tag p b)
stdDiff :: forall {k} p b (tag :: k).
(Subtractive p, ExpField b) =>
Diff tag p b -> Process (Diff tag p b) (Diff tag p b)
stdDiff = Diff tag p b -> Process (Diff tag p b) (Diff tag p b)
forall a. ExpField a => a -> Process a a
std
data DiffProcess s a b = DiffProcess
{ forall s a b. DiffProcess s a b -> Diff' a s
dInject :: Diff' a s,
forall s a b. DiffProcess s a b -> Diff' (s, a) s
dStep :: Diff' (s, a) s,
:: Diff' s b
}
data DiffSystem s i o = DiffSystem
{ :: Diff' s i,
forall s i o. DiffSystem s i o -> Diff' (s, o) s
dsStep :: Diff' (s, o) s
}
diffSystemAsDiffProcess :: DiffSystem s i o -> s -> DiffProcess s o i
diffSystemAsDiffProcess :: forall s i o. DiffSystem s i o -> s -> DiffProcess s o i
diffSystemAsDiffProcess (DiffSystem Diff' s i
ext Diff' (s, o) s
step) s
s0 =
DiffProcess
{ dInject :: Diff' o s
dInject =
(o -> (s, s -> o)) -> Diff' o s
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff
( \o
o ->
let (s
s', s -> (s, o)
pb) = Diff' (s, o) s -> (s, o) -> (s, s -> (s, o))
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' (s, o) s
step (s
s0, o
o)
in (s
s', \s
ds -> (s, o) -> o
forall a b. (a, b) -> b
snd (s -> (s, o)
pb s
ds))
),
dStep :: Diff' (s, o) s
dStep = Diff' (s, o) s
step,
dExtract :: Diff' s i
dExtract = Diff' s i
ext
}
diffProcessAsDiffSystem :: DiffProcess s a b -> DiffSystem s b a
diffProcessAsDiffSystem :: forall s a b. DiffProcess s a b -> DiffSystem s b a
diffProcessAsDiffSystem (DiffProcess Diff' a s
_ Diff' (s, a) s
step Diff' s b
ext) = Diff' s b -> Diff' (s, a) s -> DiffSystem s b a
forall s i o. Diff' s i -> Diff' (s, o) s -> DiffSystem s i o
DiffSystem Diff' s b
ext Diff' (s, a) s
step
diffSystemScan ::
(Additive s) =>
DiffSystem s i o ->
s ->
[o] ->
([i], [i] -> [o])
diffSystemScan :: forall s i o.
Additive s =>
DiffSystem s i o -> s -> [o] -> ([i], [i] -> [o])
diffSystemScan DiffSystem s i o
sys s
s0 = DiffProcess s o i -> [o] -> ([i], [i] -> [o])
forall s a b.
Additive s =>
DiffProcess s a b -> [a] -> ([b], [b] -> [a])
diffScan (DiffSystem s i o -> s -> DiffProcess s o i
forall s i o. DiffSystem s i o -> s -> DiffProcess s o i
diffSystemAsDiffProcess DiffSystem s i o
sys s
s0)
diffSystemFold ::
(Additive s, Additive i) =>
DiffSystem s i o ->
s ->
[o] ->
(i, i -> [o])
diffSystemFold :: forall s i o.
(Additive s, Additive i) =>
DiffSystem s i o -> s -> [o] -> (i, i -> [o])
diffSystemFold DiffSystem s i o
sys s
s0 = DiffProcess s o i -> [o] -> (i, i -> [o])
forall s b a.
(Additive s, Additive b) =>
DiffProcess s a b -> [a] -> (b, b -> [a])
diffFold (DiffSystem s i o -> s -> DiffProcess s o i
forall s i o. DiffSystem s i o -> s -> DiffProcess s o i
diffSystemAsDiffProcess DiffSystem s i o
sys s
s0)
data StdState a = StdState
{ forall a. StdState a -> Averager a a
stdMa :: !(Averager a a),
forall a. StdState a -> Averager a a
stdSqMa :: !(Averager a a)
}
deriving (StdState a -> StdState a -> Bool
(StdState a -> StdState a -> Bool)
-> (StdState a -> StdState a -> Bool) -> Eq (StdState a)
forall a. Eq a => StdState a -> StdState a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => StdState a -> StdState a -> Bool
== :: StdState a -> StdState a -> Bool
$c/= :: forall a. Eq a => StdState a -> StdState a -> Bool
/= :: StdState a -> StdState a -> Bool
Eq, Int -> StdState a -> ShowS
[StdState a] -> ShowS
StdState a -> String
(Int -> StdState a -> ShowS)
-> (StdState a -> String)
-> ([StdState a] -> ShowS)
-> Show (StdState a)
forall a. Show a => Int -> StdState a -> ShowS
forall a. Show a => [StdState a] -> ShowS
forall a. Show a => StdState a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> StdState a -> ShowS
showsPrec :: Int -> StdState a -> ShowS
$cshow :: forall a. Show a => StdState a -> String
show :: StdState a -> String
$cshowList :: forall a. Show a => [StdState a] -> ShowS
showList :: [StdState a] -> ShowS
Show)
instance (Additive a) => Additive (StdState a) where
zero :: StdState a
zero = Averager a a -> Averager a a -> StdState a
forall a. Averager a a -> Averager a a -> StdState a
StdState Averager a a
forall a. Additive a => a
zero Averager a a
forall a. Additive a => a
zero
StdState Averager a a
m1 Averager a a
s1 + :: StdState a -> StdState a -> StdState a
+ StdState Averager a a
m2 Averager a a
s2 = Averager a a -> Averager a a -> StdState a
forall a. Averager a a -> Averager a a -> StdState a
StdState (Averager a a
m1 Averager a a -> Averager a a -> Averager a a
forall a. Additive a => a -> a -> a
+ Averager a a
m2) (Averager a a
s1 Averager a a -> Averager a a -> Averager a a
forall a. Additive a => a -> a -> a
+ Averager a a
s2)
instance (Subtractive a) => Subtractive (StdState a) where
negate :: StdState a -> StdState a
negate (StdState Averager a a
m Averager a a
s) = Averager a a -> Averager a a -> StdState a
forall a. Averager a a -> Averager a a -> StdState a
StdState (Averager a a -> Averager a a
forall a. Subtractive a => a -> a
negate Averager a a
m) (Averager a a -> Averager a a
forall a. Subtractive a => a -> a
negate Averager a a
s)
StdState Averager a a
m1 Averager a a
s1 - :: StdState a -> StdState a -> StdState a
- StdState Averager a a
m2 Averager a a
s2 = Averager a a -> Averager a a -> StdState a
forall a. Averager a a -> Averager a a -> StdState a
StdState (Averager a a
m1 Averager a a -> Averager a a -> Averager a a
forall a. Subtractive a => a -> a -> a
- Averager a a
m2) (Averager a a
s1 Averager a a -> Averager a a -> Averager a a
forall a. Subtractive a => a -> a -> a
- Averager a a
s2)
diffForward :: DiffProcess s a b -> [a] -> (s, [s], [s -> (s, a)], s -> a)
diffForward :: forall s a b.
DiffProcess s a b -> [a] -> (s, [s], [s -> (s, a)], s -> a)
diffForward (DiffProcess Diff' a s
inj Diff' (s, a) s
step Diff' s b
_) (a
x : [a]
xs) =
let (s
s0, s -> a
injPB) = Diff' a s -> a -> (s, s -> a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' a s
inj a
x
go :: s -> [a] -> ([s], [s -> (s, a)])
go s
_ [] = ([], [])
go s
s (a
a : [a]
as) =
let (s
s', s -> (s, a)
pb) = Diff' (s, a) s -> (s, a) -> (s, s -> (s, a))
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' (s, a) s
step (s
s, a
a)
([s]
ss, [s -> (s, a)]
pbs) = s -> [a] -> ([s], [s -> (s, a)])
go s
s' [a]
as
in (s
s' s -> [s] -> [s]
forall a. a -> [a] -> [a]
: [s]
ss, s -> (s, a)
pb (s -> (s, a)) -> [s -> (s, a)] -> [s -> (s, a)]
forall a. a -> [a] -> [a]
: [s -> (s, a)]
pbs)
([s]
statesTail, [s -> (s, a)]
stepPBs) = s -> [a] -> ([s], [s -> (s, a)])
go s
s0 [a]
xs
in (s
s0, s
s0 s -> [s] -> [s]
forall a. a -> [a] -> [a]
: [s]
statesTail, [s -> (s, a)]
stepPBs, s -> a
injPB)
diffForward DiffProcess s a b
_ [] = String -> (s, [s], [s -> (s, a)], s -> a)
forall a. HasCallStack => String -> a
error String
"diffForward: empty list"
diffBackward ::
(Additive s) =>
[s] ->
[s -> (s, a)] ->
(s -> a) ->
Diff' s b ->
[b] ->
[a]
diffBackward :: forall s a b.
Additive s =>
[s] -> [s -> (s, a)] -> (s -> a) -> Diff' s b -> [b] -> [a]
diffBackward [s]
states [s -> (s, a)]
stepPBs s -> a
injPB Diff' s b
ext [b]
dys =
let n :: Int
n = [s] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [s]
states
extPBs :: [b -> s]
extPBs = (s -> b -> s) -> [s] -> [b -> s]
forall a b. (a -> b) -> [a] -> [b]
map ((b, b -> s) -> b -> s
forall a b. (a, b) -> b
snd ((b, b -> s) -> b -> s) -> (s -> (b, b -> s)) -> s -> b -> s
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Diff' s b -> s -> (b, b -> s)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' s b
ext) [s]
states
ds0 :: s
ds0 = ([b -> s]
extPBs [b -> s] -> Int -> b -> s
forall a. HasCallStack => [a] -> Int -> a
!! (Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1)) ([b]
dys [b] -> Int -> b
forall a. HasCallStack => [a] -> Int -> a
!! (Int
n Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1))
revItems :: [(s -> (s, a), b -> s, b)]
revItems = [s -> (s, a)] -> [b -> s] -> [b] -> [(s -> (s, a), b -> s, b)]
forall a b c. [a] -> [b] -> [c] -> [(a, b, c)]
zip3 ([s -> (s, a)] -> [s -> (s, a)]
forall a. [a] -> [a]
reverse [s -> (s, a)]
stepPBs) (Int -> [b -> s] -> [b -> s]
forall a. Int -> [a] -> [a]
drop Int
1 ([b -> s] -> [b -> s]
forall a. [a] -> [a]
reverse [b -> s]
extPBs)) (Int -> [b] -> [b]
forall a. Int -> [a] -> [a]
drop Int
1 ([b] -> [b]
forall a. [a] -> [a]
reverse [b]
dys))
(s
dsFinal, [a]
das) =
((s, [a]) -> (s -> (s, a), b -> s, b) -> (s, [a]))
-> (s, [a]) -> [(s -> (s, a), b -> s, b)] -> (s, [a])
forall b a. (b -> a -> b) -> b -> [a] -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl'
( \(s
ds, [a]
acc) (s -> (s, a)
pb, b -> s
extPB, b
dy) ->
let (s
dsFuture, a
da) = s -> (s, a)
pb s
ds
dsTotal :: s
dsTotal = s
dsFuture s -> s -> s
forall a. Additive a => a -> a -> a
+ b -> s
extPB b
dy
in (s
dsTotal, a
da a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
acc)
)
(s
ds0, [])
[(s -> (s, a), b -> s, b)]
revItems
da0 :: a
da0 = s -> a
injPB s
dsFinal
in a
da0 a -> [a] -> [a]
forall a. a -> [a] -> [a]
: [a]
das
diffScan :: (Additive s) => DiffProcess s a b -> [a] -> ([b], [b] -> [a])
diffScan :: forall s a b.
Additive s =>
DiffProcess s a b -> [a] -> ([b], [b] -> [a])
diffScan DiffProcess s a b
_ [] = ([], [a] -> [b] -> [a]
forall a b. a -> b -> a
const [])
diffScan DiffProcess s a b
m [a]
xs =
let (s
_, [s]
states, [s -> (s, a)]
stepPBs, s -> a
injPB) = DiffProcess s a b -> [a] -> (s, [s], [s -> (s, a)], s -> a)
forall s a b.
DiffProcess s a b -> [a] -> (s, [s], [s -> (s, a)], s -> a)
diffForward DiffProcess s a b
m [a]
xs
ys :: [b]
ys = (s -> b) -> [s] -> [b]
forall a b. (a -> b) -> [a] -> [b]
map ((b, b -> s) -> b
forall a b. (a, b) -> a
fst ((b, b -> s) -> b) -> (s -> (b, b -> s)) -> s -> b
forall b c a. (b -> c) -> (a -> b) -> a -> c
forall {k} (cat :: k -> k -> *) (b :: k) (c :: k) (a :: k).
Category cat =>
cat b c -> cat a b -> cat a c
. Diff () s b -> s -> (b, b -> s)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff (DiffProcess s a b -> Diff () s b
forall s a b. DiffProcess s a b -> Diff' s b
dExtract DiffProcess s a b
m)) [s]
states
pullback :: [b] -> [a]
pullback = [s] -> [s -> (s, a)] -> (s -> a) -> Diff () s b -> [b] -> [a]
forall s a b.
Additive s =>
[s] -> [s -> (s, a)] -> (s -> a) -> Diff' s b -> [b] -> [a]
diffBackward [s]
states [s -> (s, a)]
stepPBs s -> a
injPB (DiffProcess s a b -> Diff () s b
forall s a b. DiffProcess s a b -> Diff' s b
dExtract DiffProcess s a b
m)
in ([b]
ys, [b] -> [a]
pullback)
diffFold ::
(Additive s, Additive b) =>
DiffProcess s a b ->
[a] ->
(b, b -> [a])
diffFold :: forall s b a.
(Additive s, Additive b) =>
DiffProcess s a b -> [a] -> (b, b -> [a])
diffFold DiffProcess s a b
m [a]
xs =
let ([b]
ys, [b] -> [a]
pullback) = DiffProcess s a b -> [a] -> ([b], [b] -> [a])
forall s a b.
Additive s =>
DiffProcess s a b -> [a] -> ([b], [b] -> [a])
diffScan DiffProcess s a b
m [a]
xs
in ([b] -> b
forall a. HasCallStack => [a] -> a
last [b]
ys, \b
dy -> [b] -> [a]
pullback (Int -> b -> [b]
forall a. Int -> a -> [a]
replicate ([b] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [b]
ys Int -> Int -> Int
forall a. Subtractive a => a -> a -> a
- Int
1) b
forall a. Additive a => a
zero [b] -> [b] -> [b]
forall a. [a] -> [a] -> [a]
++ [b
dy]))
onlineDiffProcess ::
(Subtractive b, Divisive b) =>
Diff' a b ->
Diff' b b ->
DiffProcess (Averager b b) a b
onlineDiffProcess :: forall b a.
(Subtractive b, Divisive b) =>
Diff' a b -> Diff' b b -> DiffProcess (Averager b b) a b
onlineDiffProcess Diff' a b
f Diff' b b
g = Diff' a (Averager b b)
-> Diff' (Averager b b, a) (Averager b b)
-> Diff' (Averager b b) b
-> DiffProcess (Averager b b) a b
forall s a b.
Diff' a s -> Diff' (s, a) s -> Diff' s b -> DiffProcess s a b
DiffProcess Diff' a (Averager b b)
inject Diff' (Averager b b, a) (Averager b b)
step Diff' (Averager b b) b
forall {k} {p :: k}. Diff p (Averager b b) b
extract
where
inject :: Diff' a (Averager b b)
inject = (a -> (Averager b b, Averager b b -> a)) -> Diff' a (Averager b b)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a -> (Averager b b, Averager b b -> a))
-> Diff' a (Averager b b))
-> (a -> (Averager b b, Averager b b -> a))
-> Diff' a (Averager b b)
forall a b. (a -> b) -> a -> b
$ \a
x ->
let (b
y, b -> a
pb_f) = Diff' a b -> a -> (b, b -> a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' a b
f a
x
in (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
y b
forall a. Multiplicative a => a
one, \(A b
ds b
_dc) -> b -> a
pb_f b
ds)
step :: Diff' (Averager b b, a) (Averager b b)
step = ((Averager b b, a)
-> (Averager b b, Averager b b -> (Averager b b, a)))
-> Diff' (Averager b b, a) (Averager b b)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((Averager b b, a)
-> (Averager b b, Averager b b -> (Averager b b, a)))
-> Diff' (Averager b b, a) (Averager b b))
-> ((Averager b b, a)
-> (Averager b b, Averager b b -> (Averager b b, a)))
-> Diff' (Averager b b, a) (Averager b b)
forall a b. (a -> b) -> a -> b
$ \(A b
s b
c, a
x) ->
let (b
gs, b -> b
pb_gs) = Diff' b b -> b -> (b, b -> b)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' b b
g b
s
(b
gc, b -> b
pb_gc) = Diff' b b -> b -> (b, b -> b)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' b b
g b
c
(b
fx, b -> a
pb_fx) = Diff' a b -> a -> (b, b -> a)
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' a b
f a
x
s' :: b
s' = b
gs b -> b -> b
forall a. Additive a => a -> a -> a
+ b
fx
c' :: b
c' = b
gc b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one
pb :: Averager b b -> (Averager b b, a)
pb (A b
ds' b
dc') =
let ds :: b
ds = b -> b
pb_gs b
ds'
dc :: b
dc = b -> b
pb_gc b
dc'
da :: a
da = b -> a
pb_fx b
ds'
in (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
ds b
dc, a
da)
in (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
s' b
c', Averager b b -> (Averager b b, a)
pb)
extract :: Diff p (Averager b b) b
extract = (Averager b b -> (b, b -> Averager b b)) -> Diff p (Averager b b) b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Averager b b -> (b, b -> Averager b b))
-> Diff p (Averager b b) b)
-> (Averager b b -> (b, b -> Averager b b))
-> Diff p (Averager b b) b
forall a b. (a -> b) -> a -> b
$ \(A b
s b
c) ->
let y :: b
y = b
s b -> b -> b
forall a. Divisive a => a -> a -> a
/ b
c
pb :: b -> Averager b b
pb b
dy = b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A (b
dy b -> b -> b
forall a. Divisive a => a -> a -> a
/ b
c) (-((b
s b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dy) b -> b -> b
forall a. Divisive a => a -> a -> a
/ (b
c b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c)))
in (b
y, b -> Averager b b
pb)
maDiffProcess ::
(Subtractive b, Divisive b) =>
b ->
DiffProcess (Averager b b) b b
maDiffProcess :: forall b.
(Subtractive b, Divisive b) =>
b -> DiffProcess (Averager b b) b b
maDiffProcess b
r = Diff' b b -> Diff' b b -> DiffProcess (Averager b b) b b
forall b a.
(Subtractive b, Divisive b) =>
Diff' a b -> Diff' b b -> DiffProcess (Averager b b) a b
onlineDiffProcess Diff' b b
forall a. Diff () a a
forall {k} (cat :: k -> k -> *) (a :: k). Category cat => cat a a
id ((b -> (b, b -> b)) -> Diff' b b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b -> (b, b -> b)) -> Diff' b b)
-> (b -> (b, b -> b)) -> Diff' b b
forall a b. (a -> b) -> a -> b
$ \b
s -> (b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
s, (b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
*)))
sqmaDiffProcess ::
(Subtractive b, Divisive b) =>
b ->
DiffProcess (Averager b b) b b
sqmaDiffProcess :: forall b.
(Subtractive b, Divisive b) =>
b -> DiffProcess (Averager b b) b b
sqmaDiffProcess b
r = Diff' b b -> Diff' b b -> DiffProcess (Averager b b) b b
forall b a.
(Subtractive b, Divisive b) =>
Diff' a b -> Diff' b b -> DiffProcess (Averager b b) a b
onlineDiffProcess Diff' b b
forall {k} {p :: k}. Diff p b b
square ((b -> (b, b -> b)) -> Diff' b b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b -> (b, b -> b)) -> Diff' b b)
-> (b -> (b, b -> b)) -> Diff' b b
forall a b. (a -> b) -> a -> b
$ \b
s -> (b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
s, (b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
*)))
where
square :: Diff p b b
square = (b -> (b, b -> b)) -> Diff p b b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b -> (b, b -> b)) -> Diff p b b)
-> (b -> (b, b -> b)) -> Diff p b b
forall a b. (a -> b) -> a -> b
$ \b
x -> (b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
x, \b
ds' -> (b
forall a. Multiplicative a => a
one b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one) b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
ds')
stdDiffProcess ::
(Eq b, ExpField b) =>
b ->
DiffProcess (StdState b) b b
stdDiffProcess :: forall b. (Eq b, ExpField b) => b -> DiffProcess (StdState b) b b
stdDiffProcess b
r = Diff' b (StdState b)
-> Diff' (StdState b, b) (StdState b)
-> Diff' (StdState b) b
-> DiffProcess (StdState b) b b
forall s a b.
Diff' a s -> Diff' (s, a) s -> Diff' s b -> DiffProcess s a b
DiffProcess Diff' b (StdState b)
forall {k} {p :: k}. Diff p b (StdState b)
inject Diff' (StdState b, b) (StdState b)
step Diff' (StdState b) b
forall {k} {p :: k}. Diff p (StdState b) b
extract
where
inject :: Diff p b (StdState b)
inject = (b -> (StdState b, StdState b -> b)) -> Diff p b (StdState b)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((b -> (StdState b, StdState b -> b)) -> Diff p b (StdState b))
-> (b -> (StdState b, StdState b -> b)) -> Diff p b (StdState b)
forall a b. (a -> b) -> a -> b
$ \b
x ->
let pb :: StdState b -> b
pb (StdState (A b
ds b
_dc) (A b
dss b
_dcc)) =
b
ds b -> b -> b
forall a. Additive a => a -> a -> a
+ (b
forall a. Multiplicative a => a
one b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one) b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dss
in (Averager b b -> Averager b b -> StdState b
forall a. Averager a a -> Averager a a -> StdState a
StdState (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
x b
forall a. Multiplicative a => a
one) (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A (b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
x) b
forall a. Multiplicative a => a
one), StdState b -> b
pb)
step :: Diff' (StdState b, b) (StdState b)
step = ((StdState b, b) -> (StdState b, StdState b -> (StdState b, b)))
-> Diff' (StdState b, b) (StdState b)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((StdState b, b) -> (StdState b, StdState b -> (StdState b, b)))
-> Diff' (StdState b, b) (StdState b))
-> ((StdState b, b) -> (StdState b, StdState b -> (StdState b, b)))
-> Diff' (StdState b, b) (StdState b)
forall a b. (a -> b) -> a -> b
$ \(StdState (A b
s b
c) (A b
ss b
cc), b
x) ->
let s' :: b
s' = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
s b -> b -> b
forall a. Additive a => a -> a -> a
+ b
x
c' :: b
c' = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one
ss' :: b
ss' = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
ss b -> b -> b
forall a. Additive a => a -> a -> a
+ b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
x
cc' :: b
cc' = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
cc b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one
pb :: StdState b -> (StdState b, b)
pb (StdState (A b
ds' b
dc') (A b
dss' b
dcc')) =
let ds :: b
ds = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
ds'
dc :: b
dc = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dc'
dss :: b
dss = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dss'
dcc :: b
dcc = b
r b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dcc'
dx :: b
dx = b
ds' b -> b -> b
forall a. Additive a => a -> a -> a
+ (b
forall a. Multiplicative a => a
one b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one) b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
x b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dss'
in (Averager b b -> Averager b b -> StdState b
forall a. Averager a a -> Averager a a -> StdState a
StdState (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
ds b
dc) (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
dss b
dcc), b
dx)
in (Averager b b -> Averager b b -> StdState b
forall a. Averager a a -> Averager a a -> StdState a
StdState (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
s' b
c') (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
ss' b
cc'), StdState b -> (StdState b, b)
pb)
extract :: Diff p (StdState b) b
extract = (StdState b -> (b, b -> StdState b)) -> Diff p (StdState b) b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((StdState b -> (b, b -> StdState b)) -> Diff p (StdState b) b)
-> (StdState b -> (b, b -> StdState b)) -> Diff p (StdState b) b
forall a b. (a -> b) -> a -> b
$ \(StdState (A b
s b
c) (A b
ss b
cc)) ->
let y :: b
y = b -> b
forall a. ExpField a => a -> a
sqrt (b
ss b -> b -> b
forall a. Divisive a => a -> a -> a
/ b
cc b -> b -> b
forall a. Subtractive a => a -> a -> a
- (b
s b -> b -> b
forall a. Divisive a => a -> a -> a
/ b
c) b -> b -> b
forall a. Multiplicative a => a -> a -> a
* (b
s b -> b -> b
forall a. Divisive a => a -> a -> a
/ b
c))
pb :: b -> StdState b
pb b
dy
| b
y b -> b -> Bool
forall a. Eq a => a -> a -> Bool
== b
forall a. Additive a => a
zero = StdState b
forall a. Additive a => a
zero
| Bool
otherwise =
let twoY :: b
twoY = (b
forall a. Multiplicative a => a
one b -> b -> b
forall a. Additive a => a -> a -> a
+ b
forall a. Multiplicative a => a
one) b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
y
dss :: b
dss = b
dy b -> b -> b
forall a. Divisive a => a -> a -> a
/ (b
twoY b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
cc)
dcc :: b
dcc = -((b
ss b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dy) b -> b -> b
forall a. Divisive a => a -> a -> a
/ (b
twoY b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
cc b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
cc))
ds :: b
ds = -((b
s b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dy) b -> b -> b
forall a. Divisive a => a -> a -> a
/ (b
y b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c))
dc :: b
dc = (b
s b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
s b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
dy) b -> b -> b
forall a. Divisive a => a -> a -> a
/ (b
y b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c b -> b -> b
forall a. Multiplicative a => a -> a -> a
* b
c)
in Averager b b -> Averager b b -> StdState b
forall a. Averager a a -> Averager a a -> StdState a
StdState (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
ds b
dc) (b -> b -> Averager b b
forall a b. a -> b -> Averager a b
A b
dss b
dcc)
in (b
y, b -> StdState b
pb)
data DelayState a = DelayState
{ forall a. DelayState a -> a
dsOut :: !a,
forall a. DelayState a -> a
dsPrev :: !a
}
deriving (DelayState a -> DelayState a -> Bool
(DelayState a -> DelayState a -> Bool)
-> (DelayState a -> DelayState a -> Bool) -> Eq (DelayState a)
forall a. Eq a => DelayState a -> DelayState a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => DelayState a -> DelayState a -> Bool
== :: DelayState a -> DelayState a -> Bool
$c/= :: forall a. Eq a => DelayState a -> DelayState a -> Bool
/= :: DelayState a -> DelayState a -> Bool
Eq, Int -> DelayState a -> ShowS
[DelayState a] -> ShowS
DelayState a -> String
(Int -> DelayState a -> ShowS)
-> (DelayState a -> String)
-> ([DelayState a] -> ShowS)
-> Show (DelayState a)
forall a. Show a => Int -> DelayState a -> ShowS
forall a. Show a => [DelayState a] -> ShowS
forall a. Show a => DelayState a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> DelayState a -> ShowS
showsPrec :: Int -> DelayState a -> ShowS
$cshow :: forall a. Show a => DelayState a -> String
show :: DelayState a -> String
$cshowList :: forall a. Show a => [DelayState a] -> ShowS
showList :: [DelayState a] -> ShowS
Show)
instance (Additive a) => Additive (DelayState a) where
zero :: DelayState a
zero = a -> a -> DelayState a
forall a. a -> a -> DelayState a
DelayState a
forall a. Additive a => a
zero a
forall a. Additive a => a
zero
DelayState a
o1 a
p1 + :: DelayState a -> DelayState a -> DelayState a
+ DelayState a
o2 a
p2 = a -> a -> DelayState a
forall a. a -> a -> DelayState a
DelayState (a
o1 a -> a -> a
forall a. Additive a => a -> a -> a
+ a
o2) (a
p1 a -> a -> a
forall a. Additive a => a -> a -> a
+ a
p2)
instance (Subtractive a) => Subtractive (DelayState a) where
negate :: DelayState a -> DelayState a
negate (DelayState a
o a
p) = a -> a -> DelayState a
forall a. a -> a -> DelayState a
DelayState (a -> a
forall a. Subtractive a => a -> a
negate a
o) (a -> a
forall a. Subtractive a => a -> a
negate a
p)
DelayState a
o1 a
p1 - :: DelayState a -> DelayState a -> DelayState a
- DelayState a
o2 a
p2 = a -> a -> DelayState a
forall a. a -> a -> DelayState a
DelayState (a
o1 a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
o2) (a
p1 a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
p2)
delay1Diff :: (Additive a) => a -> DiffProcess (DelayState a) a a
delay1Diff :: forall a. Additive a => a -> DiffProcess (DelayState a) a a
delay1Diff a
a0 = Diff' a (DelayState a)
-> Diff' (DelayState a, a) (DelayState a)
-> Diff' (DelayState a) a
-> DiffProcess (DelayState a) a a
forall s a b.
Diff' a s -> Diff' (s, a) s -> Diff' s b -> DiffProcess s a b
DiffProcess Diff' a (DelayState a)
inject Diff' (DelayState a, a) (DelayState a)
forall {k} {p :: k}. Diff p (DelayState a, a) (DelayState a)
step Diff' (DelayState a) a
forall {k} {p :: k}. Diff p (DelayState a) a
extract
where
inject :: Diff' a (DelayState a)
inject = (a -> (DelayState a, DelayState a -> a)) -> Diff' a (DelayState a)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a -> (DelayState a, DelayState a -> a))
-> Diff' a (DelayState a))
-> (a -> (DelayState a, DelayState a -> a))
-> Diff' a (DelayState a)
forall a b. (a -> b) -> a -> b
$ \a
a -> (a -> a -> DelayState a
forall a. a -> a -> DelayState a
DelayState a
a0 a
a, DelayState a -> a
forall a. DelayState a -> a
dsPrev)
step :: Diff p (DelayState a, a) (DelayState a)
step = ((DelayState a, a)
-> (DelayState a, DelayState a -> (DelayState a, a)))
-> Diff p (DelayState a, a) (DelayState a)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((DelayState a, a)
-> (DelayState a, DelayState a -> (DelayState a, a)))
-> Diff p (DelayState a, a) (DelayState a))
-> ((DelayState a, a)
-> (DelayState a, DelayState a -> (DelayState a, a)))
-> Diff p (DelayState a, a) (DelayState a)
forall a b. (a -> b) -> a -> b
$ \(DelayState a
_ a
prev, a
a) ->
let pb :: DelayState b -> (DelayState b, b)
pb DelayState b
ds' = (b -> b -> DelayState b
forall a. a -> a -> DelayState a
DelayState b
forall a. Additive a => a
zero (DelayState b -> b
forall a. DelayState a -> a
dsOut DelayState b
ds'), DelayState b -> b
forall a. DelayState a -> a
dsPrev DelayState b
ds')
in (a -> a -> DelayState a
forall a. a -> a -> DelayState a
DelayState a
prev a
a, DelayState a -> (DelayState a, a)
forall {b}. Additive b => DelayState b -> (DelayState b, b)
pb)
extract :: Diff p (DelayState a) a
extract = (DelayState a -> (a, a -> DelayState a)) -> Diff p (DelayState a) a
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((DelayState a -> (a, a -> DelayState a))
-> Diff p (DelayState a) a)
-> (DelayState a -> (a, a -> DelayState a))
-> Diff p (DelayState a) a
forall a b. (a -> b) -> a -> b
$ \DelayState a
ds -> (DelayState a -> a
forall a. DelayState a -> a
dsOut DelayState a
ds, (a -> a -> DelayState a
forall a. a -> a -> DelayState a
`DelayState` a
forall a. Additive a => a
zero))
data DiffState a b = DiffState
{ forall a b. DiffState a b -> b
diffOut :: !b,
forall a b. DiffState a b -> a
diffPrev :: !a
}
deriving (DiffState a b -> DiffState a b -> Bool
(DiffState a b -> DiffState a b -> Bool)
-> (DiffState a b -> DiffState a b -> Bool) -> Eq (DiffState a b)
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
forall a b. (Eq b, Eq a) => DiffState a b -> DiffState a b -> Bool
$c== :: forall a b. (Eq b, Eq a) => DiffState a b -> DiffState a b -> Bool
== :: DiffState a b -> DiffState a b -> Bool
$c/= :: forall a b. (Eq b, Eq a) => DiffState a b -> DiffState a b -> Bool
/= :: DiffState a b -> DiffState a b -> Bool
Eq, Int -> DiffState a b -> ShowS
[DiffState a b] -> ShowS
DiffState a b -> String
(Int -> DiffState a b -> ShowS)
-> (DiffState a b -> String)
-> ([DiffState a b] -> ShowS)
-> Show (DiffState a b)
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
forall a b. (Show b, Show a) => Int -> DiffState a b -> ShowS
forall a b. (Show b, Show a) => [DiffState a b] -> ShowS
forall a b. (Show b, Show a) => DiffState a b -> String
$cshowsPrec :: forall a b. (Show b, Show a) => Int -> DiffState a b -> ShowS
showsPrec :: Int -> DiffState a b -> ShowS
$cshow :: forall a b. (Show b, Show a) => DiffState a b -> String
show :: DiffState a b -> String
$cshowList :: forall a b. (Show b, Show a) => [DiffState a b] -> ShowS
showList :: [DiffState a b] -> ShowS
Show)
instance (Additive a, Additive b) => Additive (DiffState a b) where
zero :: DiffState a b
zero = b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState b
forall a. Additive a => a
zero a
forall a. Additive a => a
zero
DiffState b
o1 a
p1 + :: DiffState a b -> DiffState a b -> DiffState a b
+ DiffState b
o2 a
p2 = b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState (b
o1 b -> b -> b
forall a. Additive a => a -> a -> a
+ b
o2) (a
p1 a -> a -> a
forall a. Additive a => a -> a -> a
+ a
p2)
instance (Subtractive a, Subtractive b) => Subtractive (DiffState a b) where
negate :: DiffState a b -> DiffState a b
negate (DiffState b
o a
p) = b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState (b -> b
forall a. Subtractive a => a -> a
negate b
o) (a -> a
forall a. Subtractive a => a -> a
negate a
p)
DiffState b
o1 a
p1 - :: DiffState a b -> DiffState a b -> DiffState a b
- DiffState b
o2 a
p2 = b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState (b
o1 b -> b -> b
forall a. Subtractive a => a -> a -> a
- b
o2) (a
p1 a -> a -> a
forall a. Subtractive a => a -> a -> a
- a
p2)
diffDiff ::
(Additive a, Additive b) =>
a ->
Diff' (a, a) b ->
DiffProcess (DiffState a b) a b
diffDiff :: forall a b.
(Additive a, Additive b) =>
a -> Diff' (a, a) b -> DiffProcess (DiffState a b) a b
diffDiff a
a0 Diff' (a, a) b
f = Diff' a (DiffState a b)
-> Diff' (DiffState a b, a) (DiffState a b)
-> Diff' (DiffState a b) b
-> DiffProcess (DiffState a b) a b
forall s a b.
Diff' a s -> Diff' (s, a) s -> Diff' s b -> DiffProcess s a b
DiffProcess Diff' a (DiffState a b)
inject Diff' (DiffState a b, a) (DiffState a b)
step Diff' (DiffState a b) b
forall {k} {p :: k} {b}. Diff p (DiffState a b) b
extract
where
inject :: Diff' a (DiffState a b)
inject = (a -> (DiffState a b, DiffState a b -> a))
-> Diff' a (DiffState a b)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((a -> (DiffState a b, DiffState a b -> a))
-> Diff' a (DiffState a b))
-> (a -> (DiffState a b, DiffState a b -> a))
-> Diff' a (DiffState a b)
forall a b. (a -> b) -> a -> b
$ \a
a ->
let (b
b, b -> (a, a)
pb) = Diff' (a, a) b -> (a, a) -> (b, b -> (a, a))
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' (a, a) b
f (a
a0, a
a)
pb' :: DiffState a b -> a
pb' DiffState a b
ds' = (a, a) -> a
forall a b. (a, b) -> b
snd (b -> (a, a)
pb (DiffState a b -> b
forall a b. DiffState a b -> b
diffOut DiffState a b
ds')) a -> a -> a
forall a. Additive a => a -> a -> a
+ DiffState a b -> a
forall a b. DiffState a b -> a
diffPrev DiffState a b
ds'
in (b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState b
b a
a, DiffState a b -> a
pb')
step :: Diff' (DiffState a b, a) (DiffState a b)
step = ((DiffState a b, a)
-> (DiffState a b, DiffState a b -> (DiffState a b, a)))
-> Diff' (DiffState a b, a) (DiffState a b)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((DiffState a b, a)
-> (DiffState a b, DiffState a b -> (DiffState a b, a)))
-> Diff' (DiffState a b, a) (DiffState a b))
-> ((DiffState a b, a)
-> (DiffState a b, DiffState a b -> (DiffState a b, a)))
-> Diff' (DiffState a b, a) (DiffState a b)
forall a b. (a -> b) -> a -> b
$ \(DiffState b
_ a
prev, a
a) ->
let (b
b', b -> (a, a)
pb) = Diff' (a, a) b -> (a, a) -> (b, b -> (a, a))
forall k (p :: k) a b. Diff p a b -> a -> (b, b -> a)
runDiff Diff' (a, a) b
f (a
prev, a
a)
pb' :: DiffState a b -> (DiffState a b, a)
pb' DiffState a b
ds' =
let (a
dPrev, a
dA) = b -> (a, a)
pb (DiffState a b -> b
forall a b. DiffState a b -> b
diffOut DiffState a b
ds')
in (b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState b
forall a. Additive a => a
zero a
dPrev, a
dA a -> a -> a
forall a. Additive a => a -> a -> a
+ DiffState a b -> a
forall a b. DiffState a b -> a
diffPrev DiffState a b
ds')
in (b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
DiffState b
b' a
a, DiffState a b -> (DiffState a b, a)
pb')
extract :: Diff p (DiffState a b) b
extract = (DiffState a b -> (b, b -> DiffState a b))
-> Diff p (DiffState a b) b
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((DiffState a b -> (b, b -> DiffState a b))
-> Diff p (DiffState a b) b)
-> (DiffState a b -> (b, b -> DiffState a b))
-> Diff p (DiffState a b) b
forall a b. (a -> b) -> a -> b
$ \DiffState a b
ds -> (DiffState a b -> b
forall a b. DiffState a b -> b
diffOut DiffState a b
ds, (b -> a -> DiffState a b
forall a b. b -> a -> DiffState a b
`DiffState` a
forall a. Additive a => a
zero))