{-# LANGUAGE PatternSynonyms #-}
{-# LANGUAGE RebindableSyntax #-}
{-# OPTIONS_GHC -Wno-pattern-namespace-specifier #-}

-- | Reverse-mode AD through a 'Process' scan.
--
-- The idea is to run the mealy with 'Circuit.Diff.Diff' as the carrier.  Each
-- input element becomes a variable over the whole input list, so the final
-- output (or every scan output) carries a pullback wrt every input.
--
-- The high-level gradient runners 'gradScan' and 'gradFold' are the canonical
-- entry points.  They already use 'Circuit.Diff' directly, which is the same
-- primitive arrow that @circuits-diff@ builds on.  The lower-level
-- 'DiffProcess' / 'DiffSystem' reverse-step machinery is kept because its
-- explicit-state API (capture states and step pullbacks during the forward
-- pass, then replay them backward) does not have a direct, API-preserving
-- translation to @circuits-diff@'s 'Circuit.Diff.Backprop.linearizeAt' /
-- 'Circuit.Diff.Backprop.backprop'.  Where a direct @circuits-diff@
-- replacement is possible in the future, the oracles in @test/Oracle.hs@
-- provide regression guards.
module Circuit.Stats.Diff
  ( -- * Gradient inputs
    GradInputs (..),
    variable,
    variables,

    -- * Reverse-step differentiable Process
    DiffProcess (..),
    StdState (..),
    diffFold,
    diffScan,

    -- * Reverse-step differentiable System (no inject; state external)
    DiffSystem (..),
    diffSystemAsDiffProcess,
    diffProcessAsDiffSystem,
    diffSystemScan,
    diffSystemFold,

    -- * Gradient-aware runners
    gradFold,
    gradScan,

    -- * Net Process runners
    constant,
    netFold,
    netScan,

    -- * Differentiable online statistics
    onlineDiff,
    maDiff,
    sqmaDiff,
    stdDiff,

    -- * Reverse-step online statistics
    onlineDiffProcess,
    maDiffProcess,
    sqmaDiffProcess,
    stdDiffProcess,

    -- * Reverse-step delay / diff
    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)

-- $setup
-- >>> import Circuit.Stats qualified as PS
-- >>> import Circuit.Diff

-- | A length-indexed array used as the AD parameter for a scan.
--
-- The array stores one cotangent slot per input element.  Addition and
-- subtraction are pointwise over the underlying 'Harpie.Array.Array'.
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)

-- | Componentwise maximum shape, padding missing dimensions with the longer
-- side.
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))

-- | A single input variable: the @i@-th element of an @n@-element input.
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))
  )

-- | Turn a list of inputs into a list of differentiable variables.
--
-- Each variable selects one array element in the forward pass and scatters a
-- cotangent back to that same position in the backward pass.
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]]

-- | Fold a list through a differentiable mealy and return the final output
-- together with a pullback through the entire fold.
--
-- >>> let m = PS.asum :: PS.Process (Diff' (GradInputs Double) Double) (Diff' (GradInputs Double) Double); (y, g) = gradFold m [1,2,3] in (y, g 1)
-- (6.0,[1.0,1.0,1.0])
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)

-- | Scan a list through a differentiable mealy and return the per-step
-- outputs together with a pullback that maps a list of output cotangents
-- (one per step) to input cotangents.
--
-- >>> let m = PS.asum :: PS.Process (Diff' (GradInputs Double) Double) (Diff' (GradInputs Double) Double) in snd (gradScan m [1,2,3]) [1,1,1]
-- [3.0,2.0,1.0]
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)

-- | Lift a plain input value into the parameter space of a network mealy.
--
-- The value is treated as a constant: its pullback is always zero.
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))

-- | Fold a network mealy and return the final output together with a pullback
-- through the parameters.
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

-- | Scan a network mealy and return the per-step outputs together with a
-- pullback that maps output cotangents to parameter cotangents.
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)

-- | 'Circuit.Stats.online' with a 'Diff' carrier.
--
-- The inject and step functions now operate on differentiable values, so the
-- resulting mealy can be placed inside 'netFold' / 'netScan' and the
-- parameters inside the carrier are learned by ordinary reverse-mode AD.
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

-- | Differentiable moving average with a learnable decay parameter.
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

-- | Differentiable squared moving average with a learnable decay parameter.
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

-- | Differentiable standard deviation with a learnable decay parameter.
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

-- ---------------------------------------------------------------------------
-- Reverse-step Process
-- ---------------------------------------------------------------------------

-- TODO: The 'DiffProcess' / 'DiffSystem' API below is kept because its
-- explicit-state, capture-and-replay semantics do not map directly onto
-- @circuits-ad@'s 'Net'-based 'linearizeAt' while preserving the same
-- types.  A future slice could either (a) provide a 'DiffProcess'-to-'Net'
-- compiler, or (b) deprecate this API in favour of 'gradScan' / 'gradFold'
-- plus direct @circuits-ad@ 'Circuit.Net' construction.  The oracle suite
-- guards the current behaviour.

-- | A 'Process' machine with explicit state and differentiable
-- inject / step / extract functions.
--
-- The pullbacks are captured during the forward pass and then replayed by a
-- single backward walk, so the cost of a scan gradient is linear in the input
-- length (no closure-chain blow-up).
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,
    forall s a b. DiffProcess s a b -> Diff' s b
dExtract :: Diff' s b
  }

-- | A 'System'-packaged differentiable Moore machine: state is external.
--
-- Same dynamics as 'DiffProcess', without @inject@.  Matches
-- @System s (Mono i o)@ packaging (@i@ = output, @o@ = input):
--
-- @
-- dsExtract :: Diff' s i     -- output from state
-- dsStep    :: Diff' (s, o) s  -- next state from state and input
-- @
--
-- Pair with an initial state via 'diffSystemAsDiffProcess' to reuse
-- 'diffScan' \/ 'diffFold'.  Agents that are 'System's sit here for Diff
-- without pretending to be 'Process' first.
data DiffSystem s i o = DiffSystem
  { forall s i o. DiffSystem s i o -> Diff' s i
dsExtract :: Diff' s i,
    forall s i o. DiffSystem s i o -> Diff' (s, o) s
dsStep :: Diff' (s, o) s
  }

-- | Bake in @s0@: inject is one step from that state (same as 'systemAsProcess').
--
-- @
-- diffSystemAsDiffProcess sys s0 :: DiffProcess s o i
-- @
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
    }

-- | Drop inject: keep extract and step.  Initial state must be supplied
-- again when scanning (via 'diffSystemAsDiffProcess' or 'diffSystemScan').
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

-- | Scan a 'DiffSystem' from external @s0@ (via 'diffSystemAsDiffProcess').
--
-- >>> let sys = diffProcessAsDiffSystem (maDiffProcess 0)
-- >>> let (ys, g) = diffSystemScan sys (PS.A 0 0) [1,2,3] in (ys, g [1,1,1])
-- ([1.0,2.0,3.0],[1.0,1.0,1.0])
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)

-- | Fold a 'DiffSystem' from external @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)

-- | State for a differentiable standard deviation: a moving average together
-- with a squared moving average.
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)

-- | Forward pass: capture states and step pullbacks.
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"

-- | Backward pass: walk the captured pullbacks in reverse.
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

-- | Scan a list through a 'DiffProcess' and return the per-step outputs together
-- with a pullback that maps output cotangents to input cotangents.
--
-- >>> let (ys, g) = diffScan (maDiffProcess 0) [1,2,3] in (ys, g [1,1,1])
-- ([1.0,2.0,3.0],[1.0,1.0,1.0])
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)

-- | Fold a list through a 'DiffProcess' and return the final output together with
-- a pullback through the entire fold.
--
-- >>> let (y, g) = diffFold (maDiffProcess 0) [1,2,3] in (y, g 1)
-- (3.0,[0.0,0.0,1.0])
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]))

-- | 'Circuit.Stats.online' as a reverse-step 'DiffProcess'.
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)

-- | Differentiable moving average as a reverse-step 'DiffProcess'.
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
*)))

-- | Differentiable squared moving average as a reverse-step 'DiffProcess'.
--
-- >>> let (ys, g) = diffScan (sqmaDiffProcess 0) [1,2,3] in (ys, g [1,1,1])
-- ([1.0,4.0,9.0],[2.0,4.0,6.0])
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')

-- | Differentiable standard deviation as a reverse-step 'DiffProcess'.
--
-- The gradient is taken to be zero when the standard deviation is itself zero
-- (for example, after a single sample), because the usual sqrt derivative is
-- undefined at zero.
--
-- >>> let (ys, g) = diffScan (stdDiffProcess 0) [1,2,3] in (ys, g [1,1,1])
-- ([0.0,0.0,0.0],[0.0,0.0,0.0])
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)

-- | State for a one-step delay.
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)

-- | A one-step delay as a reverse-step 'DiffProcess'.
--
-- The initial output is @a0@; after that the machine emits the previous input.
--
-- >>> let (ys, g) = diffScan (delay1Diff 0) [1,2,3] in (ys, g [1,1,1])
-- ([0,1,2],[1,1,0])
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))

-- | State for 'diffDiff'.
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)

-- | 'Circuit.Stats.diff' as a reverse-step 'DiffProcess'.
--
-- The first output uses the supplied initial previous value @a0@ instead of
-- 'undefined', which makes the machine differentiable.  This is exactly the
-- log-return pattern used in @Anal.Returns.ret@.
--
-- >>> let f = Diff $ \(p, p') -> (log (p' / p), \dy -> (negate (dy / p), dy / p'))
-- >>> let (ys, g) = diffScan (diffDiff 1 f) [1,2,4]
-- >>> (ys, g [1,1,1])
-- ([0.0,0.6931471805599453,0.6931471805599453],[0.0,0.0,0.25])
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))