{-# LANGUAGE DataKinds #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE UnicodeSyntax #-}
{-# LANGUAGE NoImplicitPrelude #-}

-- | Skip-Para neural network expressed as a polynomial coalgebra over the
-- circuits ecosystem.
--
-- * Linear maps use 'Circuit.Mat.Dense'.
-- * Gradients are produced via 'Circuit.Diff.Param.DiffP' (reverse mode).
-- * Parameter updates run through 'Circuit.Learn.Ephemeral.sgd' 'Progress'.
-- * The whole training step is packaged as a 'Circuit.System.Coalgebra'
--   whose output direction is a learning-rate scalar.
--
-- A 'These'-based boundary type is provided for future batch scheduling;
-- batching via 'Circuit.Tensor' is left as a TODO stub.
module NetCoalgebra
  ( -- * Parameter bundle (re-exported from 'Net')
    Net.NetParams (..),
    Net.netParamsFromArrays,

    -- * Forward pass
    Net.forward,
    Net.mseLoss,

    -- * Training state
    NetState (..),
    initNetState,

    -- * Training step and coalgebra
    trainStep,
    netCoalgebra,
    trainNetCoalgebra,
    trainIdentity,

    -- * Gradient oracle helpers
    gradientsViaAD,
    referenceGradients,

    -- * Boundary type
    Net.Boundary,

    -- * Batch scheduling stub
    batchSchedule,
  )
where

import Circuit.Diff.Param (DiffP (..), runDiffP)
import Circuit.Learn.Ephemeral (Progress (..), sgd)
import Circuit.Mat.Dense (Matrix (..), matTimes, matVec)
import Circuit.Poly (Eval (..), Mono, lens)
import Circuit.System (Coalgebra (..))
import Data.Vector.Unboxed qualified as VU
import Harpie.Array (Array, arrayAs)
import Harpie.Array qualified as A
import Net qualified
import NumHask.Algebra.Additive (zero)
import NumHask.Algebra.Multiplicative (one)
import Prelude hiding (id, (.))

-- | Flatten all parameters into a single vector.
flattenParams :: Net.NetParams Double -> [Double]
flattenParams :: NetParams Double -> [Double]
flattenParams NetParams Double
p =
  [[Double]] -> [Double]
forall (t :: * -> *) a. Foldable t => t [a] -> [a]
concat
    [ Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs (Matrix Double -> Array Double
forall a. Matrix a -> Array a
unMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w1 NetParams Double
p)),
      Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b1 NetParams Double
p),
      Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs (Matrix Double -> Array Double
forall a. Matrix a -> Array a
unMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p)),
      Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b2 NetParams Double
p)
    ]

-- | Reshape a flat parameter vector according to the shapes of a reference
-- parameter bundle.
reshapeParams :: [Double] -> Net.NetParams Double -> Net.NetParams Double
reshapeParams :: [Double] -> NetParams Double -> NetParams Double
reshapeParams [Double]
flat NetParams Double
p =
  let ([Double]
w1Flat, [Double]
rest1) = Int -> [Double] -> ([Double], [Double])
forall a. Int -> [a] -> ([a], [a])
splitAt (Array Double -> Int
forall a. Array a -> Int
A.size (Matrix Double -> Array Double
forall a. Matrix a -> Array a
unMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w1 NetParams Double
p))) [Double]
flat
      ([Double]
b1Flat, [Double]
rest2) = Int -> [Double] -> ([Double], [Double])
forall a. Int -> [a] -> ([a], [a])
splitAt (Array Double -> Int
forall a. Array a -> Int
A.size (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b1 NetParams Double
p)) [Double]
rest1
      ([Double]
w2Flat, [Double]
rest3) = Int -> [Double] -> ([Double], [Double])
forall a. Int -> [a] -> ([a], [a])
splitAt (Array Double -> Int
forall a. Array a -> Int
A.size (Matrix Double -> Array Double
forall a. Matrix a -> Array a
unMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p))) [Double]
rest2
      ([Double]
b2Flat, [Double]
_) = Int -> [Double] -> ([Double], [Double])
forall a. Int -> [a] -> ([a], [a])
splitAt (Array Double -> Int
forall a. Array a -> Int
A.size (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b2 NetParams Double
p)) [Double]
rest3
      mkArr :: Array a -> t -> Array a
mkArr Array a
a t
vals = [Int] -> t -> Array a
forall t a. FromVector t a => [Int] -> t -> Array a
A.array (Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a)) t
vals
   in Net.NetParams
        { w1 :: Matrix Double
Net.w1 = Array Double -> Matrix Double
forall a. Array a -> Matrix a
Matrix (Array Double -> [Double] -> Array Double
forall {t} {a} {a}. FromVector t a => Array a -> t -> Array a
mkArr (Matrix Double -> Array Double
forall a. Matrix a -> Array a
unMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w1 NetParams Double
p)) [Double]
w1Flat),
          b1 :: Array Double
Net.b1 = Array Double -> [Double] -> Array Double
forall {t} {a} {a}. FromVector t a => Array a -> t -> Array a
mkArr (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b1 NetParams Double
p) [Double]
b1Flat,
          w2 :: Matrix Double
Net.w2 = Array Double -> Matrix Double
forall a. Array a -> Matrix a
Matrix (Array Double -> [Double] -> Array Double
forall {t} {a} {a}. FromVector t a => Array a -> t -> Array a
mkArr (Matrix Double -> Array Double
forall a. Matrix a -> Array a
unMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p)) [Double]
w2Flat),
          b2 :: Array Double
Net.b2 = Array Double -> [Double] -> Array Double
forall {t} {a} {a}. FromVector t a => Array a -> t -> Array a
mkArr (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b2 NetParams Double
p) [Double]
b2Flat
        }

-- | Row count of a dense matrix.
rows :: Matrix a -> Int
rows :: forall a. Matrix a -> Int
rows (Matrix Array a
a) = case Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array a -> Vector Int
forall a. Array a -> Vector Int
A.shape Array a
a) of [Int
r, Int
_] -> Int
r; [Int]
_ -> Int
0

-- | Matrix–vector product, returning an 'Array'.
matVecArr :: Matrix Double -> Array Double -> Array Double
matVecArr :: Matrix Double -> Array Double -> Array Double
matVecArr Matrix Double
m Array Double
v = [Int] -> [Double] -> Array Double
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [Matrix Double -> Int
forall a. Matrix a -> Int
rows Matrix Double
m] (Matrix Double -> [Double] -> [Double]
forall a. (Additive a, Multiplicative a) => Matrix a -> [a] -> [a]
matVec Matrix Double
m (Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs Array Double
v))

-- | Dense matrix transpose.
transposeMatrix :: Matrix Double -> Matrix Double
transposeMatrix :: Matrix Double -> Matrix Double
transposeMatrix (Matrix Array Double
a) =
  case Vector Int -> [Int]
forall a. Unbox a => Vector a -> [a]
VU.toList (Array Double -> Vector Int
forall a. Array a -> Vector Int
A.shape Array Double
a) of
    [Int
r, Int
c] ->
      Array Double -> Matrix Double
forall a. Array a -> Matrix a
Matrix
        ( [Int] -> ([Int] -> Double) -> Array Double
forall a. [Int] -> ([Int] -> a) -> Array a
A.tabulate
            [Int
c, Int
r]
            ( \case
                [Int
i, Int
j] -> Array Double
a Array Double -> [Int] -> Double
forall a. Array a -> [Int] -> a
A.! [Int
j, Int
i]
                [Int]
_ -> [Char] -> Double
forall a. HasCallStack => [Char] -> a
error [Char]
"transposeMatrix: expected 2-element index"
            )
        )
    [Int]
_ -> [Char] -> Matrix Double
forall a. HasCallStack => [Char] -> a
error [Char]
"transposeMatrix: expected rank-2 matrix"

-- | Outer product of two vectors as a dense matrix.
outerProduct :: Array Double -> Array Double -> Matrix Double
outerProduct :: Array Double -> Array Double -> Matrix Double
outerProduct Array Double
dy Array Double
x =
  let r :: Int
r = Array Double -> Int
forall a. Array a -> Int
A.size Array Double
dy
      c :: Int
c = Array Double -> Int
forall a. Array a -> Int
A.size Array Double
x
      dyVals :: [Double]
dyVals = Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs Array Double
dy :: [Double]
      xVals :: [Double]
xVals = Array Double -> [Double]
forall t a. FromArray t a => Array a -> t
arrayAs Array Double
x :: [Double]
      dyMat :: Matrix Double
dyMat = Array Double -> Matrix Double
forall a. Array a -> Matrix a
Matrix ([Int] -> [Double] -> Array Double
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [Int
r, Int
1] [Double]
dyVals)
      xTMat :: Matrix Double
xTMat = Array Double -> Matrix Double
forall a. Array a -> Matrix a
Matrix ([Int] -> [Double] -> Array Double
forall t a. FromVector t a => [Int] -> t -> Array a
A.array [Int
1, Int
c] [Double]
xVals)
   in Matrix Double -> Matrix Double -> Matrix Double
forall a.
(Additive a, Multiplicative a) =>
Matrix a -> Matrix a -> Matrix a
matTimes Matrix Double
dyMat Matrix Double
xTMat

-- | Rectified linear unit.
reluArr :: Array Double -> Array Double
reluArr :: Array Double -> Array Double
reluArr = (Double -> Double) -> Array Double -> Array Double
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (\Double
x -> Double -> Double -> Double
forall a. Ord a => a -> a -> a
max Double
x Double
forall a. Additive a => a
zero)

-- | ReLU derivative.
reluGradArr :: Array Double -> Array Double
reluGradArr :: Array Double -> Array Double
reluGradArr = (Double -> Double) -> Array Double -> Array Double
forall a b. (a -> b) -> Array a -> Array b
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
fmap (\Double
x -> if Double
x Double -> Double -> Bool
forall a. Ord a => a -> a -> Bool
> Double
forall a. Additive a => a
zero then Double
forall a. Multiplicative a => a
one else Double
forall a. Additive a => a
zero)

-- | Full network as a parameterised differentiable arrow.
--
-- The backward pass closes over the forward intermediates and returns the
-- input cotangent together with a full 'NetParams' gradient.  This is the
-- honest 'circuits-ad' reverse-mode shape for the network.
modelDiffP :: DiffP (Net.NetParams Double) (Array Double) (Array Double)
modelDiffP :: DiffP (NetParams Double) (Array Double) (Array Double)
modelDiffP = (NetParams Double
 -> Array Double
 -> (Array Double,
     Array Double -> (Array Double, NetParams Double)))
-> DiffP (NetParams Double) (Array Double) (Array Double)
forall p a b. (p -> a -> (b, b -> (a, p))) -> DiffP p a b
DiffP ((NetParams Double
  -> Array Double
  -> (Array Double,
      Array Double -> (Array Double, NetParams Double)))
 -> DiffP (NetParams Double) (Array Double) (Array Double))
-> (NetParams Double
    -> Array Double
    -> (Array Double,
        Array Double -> (Array Double, NetParams Double)))
-> DiffP (NetParams Double) (Array Double) (Array Double)
forall a b. (a -> b) -> a -> b
$ \NetParams Double
p Array Double
x ->
  let a1 :: Array Double
a1 = Matrix Double -> Array Double -> Array Double
matVecArr (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w1 NetParams Double
p) Array Double
x
      z1 :: Array Double
z1 = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) Array Double
a1 (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b1 NetParams Double
p)
      h1 :: Array Double
h1 = Array Double -> Array Double
reluArr Array Double
z1
      a2 :: Array Double
a2 = Matrix Double -> Array Double -> Array Double
matVecArr (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p) Array Double
h1
      y :: Array Double
y = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) Array Double
a2 (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b2 NetParams Double
p)
   in ( Array Double
y,
        \Array Double
dy ->
          let dz2 :: Array Double
dz2 = Array Double
dy
              da2 :: Array Double
da2 = Array Double
dz2
              dh1 :: Array Double
dh1 = Matrix Double -> Array Double -> Array Double
matVecArr (Matrix Double -> Matrix Double
transposeMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p)) Array Double
da2
              dz1 :: Array Double
dz1 = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
dh1 (Array Double -> Array Double
reluGradArr Array Double
z1)
              da1 :: Array Double
da1 = Array Double
dz1
              dx :: Array Double
dx = Matrix Double -> Array Double -> Array Double
matVecArr (Matrix Double -> Matrix Double
transposeMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w1 NetParams Double
p)) Array Double
da1
              dw1 :: Matrix Double
dw1 = Array Double -> Array Double -> Matrix Double
outerProduct Array Double
da1 Array Double
x
              db1 :: Array Double
db1 = Array Double
da1
              dw2 :: Matrix Double
dw2 = Array Double -> Array Double -> Matrix Double
outerProduct Array Double
da2 Array Double
h1
              db2 :: Array Double
db2 = Array Double
dz2
           in ( Array Double
dx,
                Net.NetParams
                  { w1 :: Matrix Double
Net.w1 = Matrix Double
dw1,
                    b1 :: Array Double
Net.b1 = Array Double
db1,
                    w2 :: Matrix Double
Net.w2 = Matrix Double
dw2,
                    b2 :: Array Double
Net.b2 = Array Double
db2
                  }
              )
      )

-- | Training state: the parameter bundle.  The optimiser is stateless
-- plain SGD, represented as a 'Circuit.Learn.Ephemeral.Progress' that is
-- instantiated at step time with the current gradient and learning rate.
newtype NetState = NetState
  { NetState -> NetParams Double
netParams :: Net.NetParams Double
  }

-- | Initial training state.
initNetState :: Net.NetParams Double -> NetState
initNetState :: NetParams Double -> NetState
initNetState = NetParams Double -> NetState
NetState

-- | One training step.
--
-- Returns the prediction and a function that consumes the learning-rate
-- direction and produces the next state.
trainStep ::
  NetState ->
  (Array Double, Array Double) ->
  (Array Double, Double -> NetState)
trainStep :: NetState
-> (Array Double, Array Double)
-> (Array Double, Double -> NetState)
trainStep NetState
st (Array Double
x, Array Double
target) =
  let p :: NetParams Double
p = NetState -> NetParams Double
netParams NetState
st
      (Array Double
y, Array Double -> (Array Double, NetParams Double)
back) = DiffP (NetParams Double) (Array Double) (Array Double)
-> NetParams Double
-> Array Double
-> (Array Double, Array Double -> (Array Double, NetParams Double))
forall p a b. DiffP p a b -> p -> a -> (b, b -> (a, p))
runDiffP DiffP (NetParams Double) (Array Double) (Array Double)
modelDiffP NetParams Double
p Array Double
x
      (Double
_, Array Double
grad) = Array Double -> Array Double -> (Double, Array Double)
forall a.
(Fractional a, Additive a) =>
Array a -> Array a -> (a, Array a)
Net.mseLoss Array Double
y Array Double
target
      (Array Double
_dx, NetParams Double
gradParams) = Array Double -> (Array Double, NetParams Double)
back Array Double
grad
      gradFlat :: [Double]
gradFlat = NetParams Double -> [Double]
flattenParams NetParams Double
gradParams
      apply :: Double -> NetState
apply Double
lr =
        let flat :: [Double]
flat = NetParams Double -> [Double]
flattenParams NetParams Double
p
            -- 'Circuit.Learn.Ephemeral.sgd' expects a gradient function; we
            -- supply the precomputed gradient for this fixed example.
            progress :: Progress [Double] (Array Double, Array Double)
progress = Double
-> ([Double] -> (Array Double, Array Double) -> [Double])
-> Progress [Double] (Array Double, Array Double)
forall e.
Double -> ([Double] -> e -> [Double]) -> Progress [Double] e
sgd Double
lr (\[Double]
_ (Array Double, Array Double)
_ -> [Double]
gradFlat)
            flat' :: [Double]
flat' = Progress [Double] (Array Double, Array Double)
-> (Array Double, Array Double) -> [Double] -> [Double]
forall p e. Progress p e -> e -> p -> p
step Progress [Double] (Array Double, Array Double)
progress (Array Double
x, Array Double
target) [Double]
flat
         in NetParams Double -> NetState
NetState ([Double] -> NetParams Double -> NetParams Double
reshapeParams [Double]
flat' NetParams Double
p)
   in (Array Double
y, Double -> NetState
apply)

-- | Polynomial coalgebra view of the network.
--
-- * State: 'NetState'.
-- * Input position: @(input, target)@; no input direction.
-- * Output position: prediction; output direction: learning rate.
netCoalgebra ::
  Coalgebra
    NetState
    (Mono () (Array Double, Array Double))
    (Mono Double (Array Double))
netCoalgebra :: Coalgebra
  NetState
  (Mono () (Array Double, Array Double))
  (Mono Double (Array Double))
netCoalgebra =
  Coalgebra
    { act :: NetState
-> Morphism
     (Mono () (Array Double, Array Double)) (Mono Double (Array Double))
act =
        \NetState
st ->
          ((Array Double, Array Double) -> Array Double)
-> ((Array Double, Array Double) -> Double -> ())
-> Morphism
     (Mono () (Array Double, Array Double)) (Mono Double (Array Double))
forall a b db da.
(a -> b) -> (a -> db -> da) -> Morphism (Mono da a) (Mono db b)
lens
            (\(Array Double
x, Array Double
_) -> NetParams Double -> Array Double -> Array Double
forall a.
(Ord a, Num a, Additive a, Multiplicative a) =>
NetParams a -> Array a -> Array a
Net.forward (NetState -> NetParams Double
netParams NetState
st) Array Double
x)
            (\(Array Double, Array Double)
_ Double
_ -> ()),
      upd :: NetState
-> Eval (Mono () (Array Double, Array Double)) NetState
-> Step NetState (Mono Double (Array Double))
upd =
        \NetState
st (EP (EK (Array Double
x, Array Double
target), Eval q NetState
_)) ->
          let (Array Double
y, Double -> NetState
apply) = NetState
-> (Array Double, Array Double)
-> (Array Double, Double -> NetState)
trainStep NetState
st (Array Double
x, Array Double
target)
           in (Eval ('Const (Array Double)) NetState,
 Eval ('Exp Double) NetState)
-> Step NetState (Mono Double (Array Double))
forall (p1 :: Poly) x (q :: Poly).
(Eval p1 x, Eval q x) -> Eval ('Prod p1 q) x
EP (Array Double -> Eval ('Const (Array Double)) NetState
forall c x. c -> Eval ('Const c) x
EK Array Double
y, (Double -> NetState) -> Eval ('Exp Double) NetState
forall a x. (a -> x) -> Eval ('Exp a) x
EE Double -> NetState
apply)
    }

-- | Train the coalgebra for @n@ steps on a fixed @(x, target)@ pair.
trainNetCoalgebra ::
  Net.NetParams Double ->
  Array Double ->
  Array Double ->
  Double ->
  Int ->
  [(Double, Net.NetParams Double)]
trainNetCoalgebra :: NetParams Double
-> Array Double
-> Array Double
-> Double
-> Int
-> [(Double, NetParams Double)]
trainNetCoalgebra NetParams Double
p0 Array Double
x Array Double
target Double
lr Int
n = Int
-> Double
-> NetParams Double
-> Array Double
-> Array Double
-> [(Double, NetParams Double)]
trainIdentity Int
n Double
lr NetParams Double
p0 Array Double
x Array Double
target

-- | Identity-mapping training oracle: train for @n@ steps and return the
-- loss before each update together with the parameters after each update.
trainIdentity ::
  Int ->
  Double ->
  Net.NetParams Double ->
  Array Double ->
  Array Double ->
  [(Double, Net.NetParams Double)]
trainIdentity :: Int
-> Double
-> NetParams Double
-> Array Double
-> Array Double
-> [(Double, NetParams Double)]
trainIdentity Int
n Double
lr NetParams Double
p0 Array Double
x Array Double
target = NetState -> Int -> [(Double, NetParams Double)]
go (NetParams Double -> NetState
initNetState NetParams Double
p0) Int
n
  where
    go :: NetState -> Int -> [(Double, NetParams Double)]
go NetState
_ Int
0 = []
    go NetState
st Int
k =
      let y :: Array Double
y = NetParams Double -> Array Double -> Array Double
forall a.
(Ord a, Num a, Additive a, Multiplicative a) =>
NetParams a -> Array a -> Array a
Net.forward (NetState -> NetParams Double
netParams NetState
st) Array Double
x
          (Double
loss, Array Double
_) = Array Double -> Array Double -> (Double, Array Double)
forall a.
(Fractional a, Additive a) =>
Array a -> Array a -> (a, Array a)
Net.mseLoss Array Double
y Array Double
target
          (Array Double
_, Double -> NetState
apply) = NetState
-> (Array Double, Array Double)
-> (Array Double, Double -> NetState)
trainStep NetState
st (Array Double
x, Array Double
target)
          st' :: NetState
st' = Double -> NetState
apply Double
lr
       in (Double
loss, NetState -> NetParams Double
netParams NetState
st') (Double, NetParams Double)
-> [(Double, NetParams Double)] -> [(Double, NetParams Double)]
forall a. a -> [a] -> [a]
: NetState -> Int -> [(Double, NetParams Double)]
go NetState
st' (Int
k Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1)

-- | Compute parameter gradients via the 'circuits-ad' 'DiffP' model.
gradientsViaAD ::
  Net.NetParams Double ->
  Array Double ->
  Array Double ->
  Net.NetParams Double
gradientsViaAD :: NetParams Double
-> Array Double -> Array Double -> NetParams Double
gradientsViaAD NetParams Double
p Array Double
x Array Double
target =
  let (Array Double
y, Array Double -> (Array Double, NetParams Double)
back) = DiffP (NetParams Double) (Array Double) (Array Double)
-> NetParams Double
-> Array Double
-> (Array Double, Array Double -> (Array Double, NetParams Double))
forall p a b. DiffP p a b -> p -> a -> (b, b -> (a, p))
runDiffP DiffP (NetParams Double) (Array Double) (Array Double)
modelDiffP NetParams Double
p Array Double
x
      (Double
_, Array Double
grad) = Array Double -> Array Double -> (Double, Array Double)
forall a.
(Fractional a, Additive a) =>
Array a -> Array a -> (a, Array a)
Net.mseLoss Array Double
y Array Double
target
      (Array Double
_, NetParams Double
gradParams) = Array Double -> (Array Double, NetParams Double)
back Array Double
grad
   in NetParams Double
gradParams

-- | Hand-rolled reference gradients for the 2-layer network, kept to verify
-- the 'circuits-ad' / 'DiffP' gradients in the oracle.
--
-- This is the old skip-Para chain rule, written directly against the same
-- 'Circuit.Mat.Dense' primitives used by the rest of the module.
referenceGradients ::
  Net.NetParams Double ->
  Array Double ->
  Array Double ->
  Net.NetParams Double
referenceGradients :: NetParams Double
-> Array Double -> Array Double -> NetParams Double
referenceGradients NetParams Double
p Array Double
x Array Double
target =
  let a1 :: Array Double
a1 = Matrix Double -> Array Double -> Array Double
matVecArr (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w1 NetParams Double
p) Array Double
x
      z1 :: Array Double
z1 = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) Array Double
a1 (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b1 NetParams Double
p)
      h1 :: Array Double
h1 = Array Double -> Array Double
reluArr Array Double
z1
      a2 :: Array Double
a2 = Matrix Double -> Array Double -> Array Double
matVecArr (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p) Array Double
h1
      y :: Array Double
y = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) Array Double
a2 (NetParams Double -> Array Double
forall a. NetParams a -> Array a
Net.b2 NetParams Double
p)
      dOut :: Array Double
dOut = (Double, Array Double) -> Array Double
forall a b. (a, b) -> b
snd (Array Double -> Array Double -> (Double, Array Double)
forall a.
(Fractional a, Additive a) =>
Array a -> Array a -> (a, Array a)
Net.mseLoss Array Double
y Array Double
target)
      db2 :: Array Double
db2 = Array Double
dOut
      dw2 :: Matrix Double
dw2 = Array Double -> Array Double -> Matrix Double
outerProduct Array Double
dOut Array Double
h1
      da2 :: Array Double
da2 = Array Double
dOut
      dh1 :: Array Double
dh1 = Matrix Double -> Array Double -> Array Double
matVecArr (Matrix Double -> Matrix Double
transposeMatrix (NetParams Double -> Matrix Double
forall a. NetParams a -> Matrix a
Net.w2 NetParams Double
p)) Array Double
da2
      dz1 :: Array Double
dz1 = (Double -> Double -> Double)
-> Array Double -> Array Double -> Array Double
forall a b c. (a -> b -> c) -> Array a -> Array b -> Array c
A.zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(*) Array Double
dh1 (Array Double -> Array Double
reluGradArr Array Double
z1)
      db1 :: Array Double
db1 = Array Double
dz1
      da1 :: Array Double
da1 = Array Double
dz1
      dw1 :: Matrix Double
dw1 = Array Double -> Array Double -> Matrix Double
outerProduct Array Double
da1 Array Double
x
   in Matrix Double
-> Array Double
-> Matrix Double
-> Array Double
-> NetParams Double
forall a. Matrix a -> Array a -> Matrix a -> Array a -> NetParams a
Net.NetParams Matrix Double
dw1 Array Double
db1 Matrix Double
dw2 Array Double
db2

-- | Batch scheduling stub.  The intended implementation uses
-- 'Circuit.Tensor' combinators (broadcast / telecast / indexWindows) to
-- turn single-example layers into batch layers and fold gradients across
-- the batch dimension.
batchSchedule :: a
batchSchedule :: forall a. a
batchSchedule = [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"batchSchedule: not yet implemented"