{-# LANGUAGE DataKinds #-}
{-# LANGUAGE GADTs #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# LANGUAGE UnicodeSyntax #-}
{-# LANGUAGE NoImplicitPrelude #-}
module NetCoalgebra
(
Net.NetParams (..),
Net.netParamsFromArrays,
Net.forward,
Net.mseLoss,
NetState (..),
initNetState,
trainStep,
netCoalgebra,
trainNetCoalgebra,
trainIdentity,
gradientsViaAD,
referenceGradients,
Net.Boundary,
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, (.))
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)
]
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
}
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
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))
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"
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
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)
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)
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
}
)
)
newtype NetState = NetState
{ NetState -> NetParams Double
netParams :: Net.NetParams Double
}
initNetState :: Net.NetParams Double -> NetState
initNetState :: NetParams Double -> NetState
initNetState = NetParams Double -> NetState
NetState
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
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)
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)
}
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
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)
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
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
batchSchedule :: a
batchSchedule :: forall a. a
batchSchedule = [Char] -> a
forall a. HasCallStack => [Char] -> a
error [Char]
"batchSchedule: not yet implemented"