{-# LANGUAGE DataKinds #-}
{-# LANGUAGE RebindableSyntax #-}
module Circuit.Diff.Array
(
tensorTaylor,
tensorDerivativeN,
elementwiseTower,
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
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))
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
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
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)
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)
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)
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)
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)
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))
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))
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
)