{-# LANGUAGE DataKinds #-}
{-# LANGUAGE RebindableSyntax #-}

-- | Array AD on top of 'Harpie.Fixed.Array' and 'Circuit.Mat.Square'.
--
-- This module provides:
--
-- * forward-mode Taylor towers for square-matrix computations, via 'Jet'
--   instantiated at 'Square n Double';
-- * elementwise towers by mapping scalar 'Jet's across array elements;
-- * explicit reverse-mode 'Diff' primitives for matrix multiplication,
--   transpose, sum, scale, and elementwise activations.
--
-- The generic 'Multiplicative' instance of 'Diff' is /not/ used for matrix
-- multiplication because its product-rule pullback is the wrong adjoint for
-- non-commutative multiplication.
module Circuit.Diff.Array
  ( -- * Forward-mode tensor towers
    tensorTaylor,
    tensorDerivativeN,
    elementwiseTower,

    -- * Reverse-mode tensor primitives
    matMulD,
    transposeD,
    sumD,
    scaleD,
    elementwiseD,
    sigmoidD,
    tanhD,
  )
where

import Circuit.Diff (Diff (..))
import Circuit.Diff.Jet (Jet (..), taylorDers, variable)
import Circuit.Mat.Square (Square)
import Data.Foldable (foldl')
import Data.Proxy (Proxy (..))
import GHC.TypeNats (KnownNat)
import Harpie.Fixed (Array)
import Harpie.Fixed qualified as F
import Harpie.Shape (KnownNats)
import NumHask.Algebra.Additive (Additive (..), Subtractive (..))
import NumHask.Algebra.Field (ExpField (..), TrigField (..))
import NumHask.Algebra.Multiplicative (Divisive (..), Multiplicative (..), recip)
import NumHask.Data.Integral (FromInteger (..))
import NumHask.Prelude

-- ---------------------------------------------------------------------------
-- Forward-mode tensor towers
-- ---------------------------------------------------------------------------

-- | First @k@ raw derivatives of a square-matrix function at @x0@.
--
-- The function must be built from NumHask-polymorphic operations so that the
-- 'Jet' recurrence engine can propagate matrix coefficients through matrix
-- multiplication.
tensorTaylor ::
  forall n a.
  ( KnownNat n,
    Additive a,
    Multiplicative a,
    FromInteger a
  ) =>
  (Jet (Square n a) -> Jet (Square n a)) ->
  Int ->
  Square n a ->
  [Square n a]
tensorTaylor :: forall (n :: Nat) a.
(KnownNat n, Additive a, Multiplicative a, FromInteger a) =>
(Jet (Square n a) -> Jet (Square n a))
-> Int -> Square n a -> [Square n a]
tensorTaylor Jet (Square n a) -> Jet (Square n a)
f Int
k Square n a
x0 = Jet (Square n a) -> [Square n a]
forall a.
(Additive a, Multiplicative a, FromInteger a) =>
Jet a -> [a]
taylorDers (Jet (Square n a) -> Jet (Square n a)
f (Int -> Square n a -> Jet (Square n a)
forall a. (Additive a, Multiplicative a) => Int -> a -> Jet a
variable Int
k Square n a
x0))

-- | The @n@-th raw derivative of a square-matrix function at @x0@.
tensorDerivativeN ::
  forall n a.
  ( KnownNat n,
    Additive a,
    Multiplicative a,
    FromInteger a
  ) =>
  (Jet (Square n a) -> Jet (Square n a)) ->
  Int ->
  Square n a ->
  Square n a
tensorDerivativeN :: forall (n :: Nat) a.
(KnownNat n, Additive a, Multiplicative a, FromInteger a) =>
(Jet (Square n a) -> Jet (Square n a))
-> Int -> Square n a -> Square n a
tensorDerivativeN Jet (Square n a) -> Jet (Square n a)
f Int
n Square n a
x0 = (Jet (Square n a) -> Jet (Square n a))
-> Int -> Square n a -> [Square n a]
forall (n :: Nat) a.
(KnownNat n, Additive a, Multiplicative a, FromInteger a) =>
(Jet (Square n a) -> Jet (Square n a))
-> Int -> Square n a -> [Square n a]
tensorTaylor Jet (Square n a) -> Jet (Square n a)
f Int
n Square n a
x0 [Square n a] -> Int -> Square n a
forall a. HasCallStack => [a] -> Int -> a
!! Int
n

-- | Apply a scalar tower function elementwise to every entry of an array.
--
-- This is the elementwise analogue of Eshkol's @tt-sigmoid@ / @tt-tanh@:
-- each component gets its own scalar Taylor tower, independent of the others.
elementwiseTower ::
  forall s a.
  ( KnownNats s,
    Additive a,
    Multiplicative a,
    FromInteger a
  ) =>
  (Jet a -> Jet a) ->
  Int ->
  Array s a ->
  Array s a
elementwiseTower :: forall (s :: [Nat]) a.
(KnownNats s, Additive a, Multiplicative a, FromInteger a) =>
(Jet a -> Jet a) -> Int -> Array s a -> Array s a
elementwiseTower Jet a -> Jet a
f Int
k Array s a
arr =
  (Rep (Array s) -> a) -> Array s a
forall a. (Rep (Array s) -> a) -> Array s a
forall (f :: * -> *) a. Representable f => (Rep f -> a) -> f a
F.tabulate ((Rep (Array s) -> a) -> Array s a)
-> (Rep (Array s) -> a) -> Array s a
forall a b. (a -> b) -> a -> b
$ \Rep (Array s)
ix ->
    let x :: a
x = Array s a -> Rep (Array s) -> a
forall a. Array s a -> Rep (Array s) -> a
forall (f :: * -> *) a. Representable f => f a -> Rep f -> a
F.index Array s a
arr Rep (Array s)
ix
     in Jet a -> [a]
forall a.
(Additive a, Multiplicative a, FromInteger a) =>
Jet a -> [a]
taylorDers (Jet a -> Jet a
f (Int -> a -> Jet a
forall a. (Additive a, Multiplicative a) => Int -> a -> Jet a
variable Int
k a
x)) [a] -> Int -> a
forall a. HasCallStack => [a] -> Int -> a
!! Int
k

-- ---------------------------------------------------------------------------
-- Reverse-mode tensor primitives
-- ---------------------------------------------------------------------------

-- | Matrix multiplication with the correct reverse-mode adjoint.
--
-- For @Y = A B@ the pullback is @dA = dY B^T@, @dB = A^T dY@.
matMulD ::
  forall n p.
  (KnownNat n) =>
  Diff p (Square n Double, Square n Double) (Square n Double)
matMulD :: forall {k} (n :: Nat) (p :: k).
KnownNat n =>
Diff p (Square n Double, Square n Double) (Square n Double)
matMulD =
  ((Array '[n, n] Double, Array '[n, n] Double)
 -> (Array '[n, n] Double,
     Array '[n, n] Double
     -> (Array '[n, n] Double, Array '[n, n] Double)))
-> Diff
     p
     (Array '[n, n] Double, Array '[n, n] Double)
     (Array '[n, n] Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff (((Array '[n, n] Double, Array '[n, n] Double)
  -> (Array '[n, n] Double,
      Array '[n, n] Double
      -> (Array '[n, n] Double, Array '[n, n] Double)))
 -> Diff
      p
      (Array '[n, n] Double, Array '[n, n] Double)
      (Array '[n, n] Double))
-> ((Array '[n, n] Double, Array '[n, n] Double)
    -> (Array '[n, n] Double,
        Array '[n, n] Double
        -> (Array '[n, n] Double, Array '[n, n] Double)))
-> Diff
     p
     (Array '[n, n] Double, Array '[n, n] Double)
     (Array '[n, n] Double)
forall a b. (a -> b) -> a -> b
$ \(Array '[n, n] Double
a, Array '[n, n] Double
b) ->
    let y :: Array '[n, n] Double
y = Array '[n, n] Double
a Array '[n, n] Double
-> Array '[n, n] Double -> Array '[n, n] Double
forall a. Multiplicative a => a -> a -> a
* Array '[n, n] Double
b
        bt :: Array '[n, n] Double
bt = Array '[n, n] Double -> Array '[n, n] Double
forall a (s :: [Nat]) (s' :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (Reverse s)) =>
Array s a -> Array s' a
F.transpose Array '[n, n] Double
b
        at :: Array '[n, n] Double
at = Array '[n, n] Double -> Array '[n, n] Double
forall a (s :: [Nat]) (s' :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (Reverse s)) =>
Array s a -> Array s' a
F.transpose Array '[n, n] Double
a
        pb :: Array '[n, n] Double
-> (Array '[n, n] Double, Array '[n, n] Double)
pb Array '[n, n] Double
db = (Array '[n, n] Double
db Array '[n, n] Double
-> Array '[n, n] Double -> Array '[n, n] Double
forall a. Multiplicative a => a -> a -> a
* Array '[n, n] Double
bt, Array '[n, n] Double
at Array '[n, n] Double
-> Array '[n, n] Double -> Array '[n, n] Double
forall a. Multiplicative a => a -> a -> a
* Array '[n, n] Double
db)
     in (Array '[n, n] Double
y, Array '[n, n] Double
-> (Array '[n, n] Double, Array '[n, n] Double)
pb)

-- | Transpose.  Its own adjoint.
transposeD ::
  forall m n p.
  (KnownNat m, KnownNat n) =>
  Diff p (Array '[m, n] Double) (Array '[n, m] Double)
transposeD :: forall {k} (m :: Nat) (n :: Nat) (p :: k).
(KnownNat m, KnownNat n) =>
Diff p (Array '[m, n] Double) (Array '[n, m] Double)
transposeD =
  (Array '[m, n] Double
 -> (Array '[n, m] Double,
     Array '[n, m] Double -> Array '[m, n] Double))
-> Diff p (Array '[m, n] Double) (Array '[n, m] Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array '[m, n] Double
  -> (Array '[n, m] Double,
      Array '[n, m] Double -> Array '[m, n] Double))
 -> Diff p (Array '[m, n] Double) (Array '[n, m] Double))
-> (Array '[m, n] Double
    -> (Array '[n, m] Double,
        Array '[n, m] Double -> Array '[m, n] Double))
-> Diff p (Array '[m, n] Double) (Array '[n, m] Double)
forall a b. (a -> b) -> a -> b
$ \Array '[m, n] Double
a ->
    let y :: Array '[n, m] Double
y = Array '[m, n] Double -> Array '[n, m] Double
forall a (s :: [Nat]) (s' :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (Reverse s)) =>
Array s a -> Array s' a
F.transpose Array '[m, n] Double
a
     in (Array '[n, m] Double
y, Array '[n, m] Double -> Array '[m, n] Double
forall a (s :: [Nat]) (s' :: [Nat]).
(KnownNats s, KnownNats s', s' ~ Eval (Reverse s)) =>
Array s a -> Array s' a
F.transpose)

-- | Sum all elements of an array.  Pullback replicates the scalar.
sumD ::
  forall s p.
  (KnownNats s) =>
  Diff p (Array s Double) Double
sumD :: forall {k} (s :: [Nat]) (p :: k).
KnownNats s =>
Diff p (Array s Double) Double
sumD =
  (Array s Double -> (Double, Double -> Array s Double))
-> Diff p (Array s Double) Double
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array s Double -> (Double, Double -> Array s Double))
 -> Diff p (Array s Double) Double)
-> (Array s Double -> (Double, Double -> Array s Double))
-> Diff p (Array s Double) Double
forall a b. (a -> b) -> a -> b
$ \Array s Double
a ->
    let y :: Double
y = (Double -> Double -> Double) -> Double -> Array s Double -> Double
forall b a. (b -> a -> b) -> b -> Array s a -> b
forall (t :: * -> *) b a.
Foldable t =>
(b -> a -> b) -> b -> t a -> b
foldl' Double -> Double -> Double
forall a. Additive a => a -> a -> a
(+) Double
forall a. Additive a => a
zero Array s Double
a
     in (Double
y, \Double
db -> Double -> Array s Double
forall (s :: [Nat]) a. KnownNats s => a -> Array s a
F.konst Double
db)

-- | Scale every element by a constant.
scaleD ::
  forall s p.
  Double ->
  Diff p (Array s Double) (Array s Double)
scaleD :: forall {k} (s :: [Nat]) (p :: k).
Double -> Diff p (Array s Double) (Array s Double)
scaleD Double
s =
  (Array s Double
 -> (Array s Double, Array s Double -> Array s Double))
-> Diff p (Array s Double) (Array s Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array s Double
  -> (Array s Double, Array s Double -> Array s Double))
 -> Diff p (Array s Double) (Array s Double))
-> (Array s Double
    -> (Array s Double, Array s Double -> Array s Double))
-> Diff p (Array s Double) (Array s Double)
forall a b. (a -> b) -> a -> b
$ \Array s Double
a ->
    let y :: Array s Double
y = (Double -> Double) -> Array s Double -> Array s Double
forall a b. (a -> b) -> Array s a -> Array s b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (Double
s Double -> Double -> Double
forall a. Multiplicative a => a -> a -> a
*) Array s Double
a
     in (Array s Double
y, \Array s Double
db -> (Double -> Double) -> Array s Double -> Array s Double
forall a b. (a -> b) -> Array s a -> Array s b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (Double
s Double -> Double -> Double
forall a. Multiplicative a => a -> a -> a
*) Array s Double
db)

-- | Elementwise activation from a scalar primitive @x -> (y, dy/dx)@.
elementwiseD ::
  forall s p.
  (KnownNats s) =>
  (Double -> (Double, Double -> Double)) ->
  Diff p (Array s Double) (Array s Double)
elementwiseD :: forall {k} (s :: [Nat]) (p :: k).
KnownNats s =>
(Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
elementwiseD Double -> (Double, Double -> Double)
phi =
  (Array s Double
 -> (Array s Double, Array s Double -> Array s Double))
-> Diff p (Array s Double) (Array s Double)
forall k (p :: k) a b. (a -> (b, b -> a)) -> Diff p a b
Diff ((Array s Double
  -> (Array s Double, Array s Double -> Array s Double))
 -> Diff p (Array s Double) (Array s Double))
-> (Array s Double
    -> (Array s Double, Array s Double -> Array s Double))
-> Diff p (Array s Double) (Array s Double)
forall a b. (a -> b) -> a -> b
$ \Array s Double
a ->
    let (Array s Double
ys, Array s (Double -> Double)
grads) = Array s (Double, Double -> Double)
-> (Array s Double, Array s (Double -> Double))
forall (s :: [Nat]) a b. Array s (a, b) -> (Array s a, Array s b)
unzipA ((Double -> (Double, Double -> Double))
-> Array s Double -> Array s (Double, Double -> Double)
forall a b. (a -> b) -> Array s a -> Array s b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap Double -> (Double, Double -> Double)
phi Array s Double
a)
        pb :: Array s Double -> Array s Double
pb Array s Double
db = ((Double -> Double) -> Double -> Double)
-> Array s (Double -> Double) -> Array s Double -> Array s Double
forall (s :: [Nat]) a b c.
KnownNats s =>
(a -> b -> c) -> Array s a -> Array s b -> Array s c
F.zipWith (\Double -> Double
g Double
dbi -> Double -> Double
g Double
dbi) Array s (Double -> Double)
grads Array s Double
db
     in (Array s Double
ys, Array s Double -> Array s Double
pb)

-- | Sigmoid activation.
sigmoidD ::
  forall s p.
  (KnownNats s) =>
  Diff p (Array s Double) (Array s Double)
sigmoidD :: forall {k} (s :: [Nat]) (p :: k).
KnownNats s =>
Diff p (Array s Double) (Array s Double)
sigmoidD = (Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
forall {k} (s :: [Nat]) (p :: k).
KnownNats s =>
(Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
elementwiseD ((Double -> (Double, Double -> Double))
 -> Diff p (Array s Double) (Array s Double))
-> (Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
forall a b. (a -> b) -> a -> b
$ \Double
x ->
  let s :: Double
s = Double -> Double
forall a. Divisive a => a -> a
recip (Double
1 Double -> Double -> Double
forall a. Additive a => a -> a -> a
+ Double -> Double
forall a. ExpField a => a -> a
exp (-Double
x))
   in (Double
s, \Double
db -> Double
db Double -> Double -> Double
forall a. Multiplicative a => a -> a -> a
* Double
s Double -> Double -> Double
forall a. Multiplicative a => a -> a -> a
* (Double
1 Double -> Double -> Double
forall a. Subtractive a => a -> a -> a
- Double
s))

-- | Tanh activation.
tanhD ::
  forall s p.
  (KnownNats s) =>
  Diff p (Array s Double) (Array s Double)
tanhD :: forall {k} (s :: [Nat]) (p :: k).
KnownNats s =>
Diff p (Array s Double) (Array s Double)
tanhD = (Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
forall {k} (s :: [Nat]) (p :: k).
KnownNats s =>
(Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
elementwiseD ((Double -> (Double, Double -> Double))
 -> Diff p (Array s Double) (Array s Double))
-> (Double -> (Double, Double -> Double))
-> Diff p (Array s Double) (Array s Double)
forall a b. (a -> b) -> a -> b
$ \Double
x ->
  let t :: Double
t = Double -> Double
forall a. TrigField a => a -> a
tanh Double
x
   in (Double
t, \Double
db -> Double
db Double -> Double -> Double
forall a. Multiplicative a => a -> a -> a
* (Double
1 Double -> Double -> Double
forall a. Subtractive a => a -> a -> a
- Double
t Double -> Double -> Double
forall a. Multiplicative a => a -> a -> a
* Double
t))

-- ---------------------------------------------------------------------------
-- Helpers
-- ---------------------------------------------------------------------------

-- Split an array of @(value, derivative-function)@ pairs into two arrays.
unzipA ::
  Array s (a, b) ->
  (Array s a, Array s b)
unzipA :: forall (s :: [Nat]) a b. Array s (a, b) -> (Array s a, Array s b)
unzipA Array s (a, b)
ab =
  ( ((a, b) -> a) -> Array s (a, b) -> Array s a
forall a b. (a -> b) -> Array s a -> Array s b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (a, b) -> a
forall a b. (a, b) -> a
fst Array s (a, b)
ab,
    ((a, b) -> b) -> Array s (a, b) -> Array s b
forall a b. (a -> b) -> Array s a -> Array s b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (a, b) -> b
forall a b. (a, b) -> b
snd Array s (a, b)
ab
  )