module Circuit.Learn.Adam
(
ewma,
ewmaDirect,
adam,
adamDecomposed,
adamReference,
adamUpdates,
)
where
import Circuit.Process (Process (..), scan)
import Data.List (scanl')
ewma :: Double -> Double -> Process Double Double
ewma :: Double -> Double -> Process Double Double
ewma Double
alpha Double
s0 = (Double -> Double)
-> (Double -> Double -> Double)
-> (Double -> Double)
-> Process Double Double
forall s a b. (a -> s) -> (s -> a -> s) -> (s -> b) -> Process a b
Process Double -> Double
inject Double -> Double -> Double
step Double -> Double
forall {p}. p -> p
extract
where
inject :: Double -> Double
inject Double
x0 = Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x0 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
alpha) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
s0
step :: Double -> Double -> Double
step Double
s Double
x = Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
alpha) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
s
extract :: p -> p
extract p
s = p
s
ewmaDirect :: Double -> Double -> [Double] -> [Double]
ewmaDirect :: Double -> Double -> [Double] -> [Double]
ewmaDirect Double
_ Double
_ [] = []
ewmaDirect Double
alpha Double
s0 (Double
x0 : [Double]
xs) = (Double -> Double -> Double) -> Double -> [Double] -> [Double]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' Double -> Double -> Double
step (Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x0 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
alpha) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
s0) [Double]
xs
where
step :: Double -> Double -> Double
step Double
s Double
x = Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
alpha) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
s
adam :: Double -> Double -> Double -> Double -> Process Double Double
adam :: Double -> Double -> Double -> Double -> Process Double Double
adam Double
alpha Double
beta1 Double
beta2 Double
eps = (Double -> (Double, Double, Double))
-> ((Double, Double, Double) -> Double -> (Double, Double, Double))
-> ((Double, Double, Double) -> Double)
-> Process Double Double
forall s a b. (a -> s) -> (s -> a -> s) -> (s -> b) -> Process a b
Process Double -> (Double, Double, Double)
inject (Double, Double, Double) -> Double -> (Double, Double, Double)
step (Double, Double, Double) -> Double
extract
where
inject :: Double -> (Double, Double, Double)
inject Double
g0 =
let m1 :: Double
m1 = (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g0
v1 :: Double
v1 = (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g0 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g0
in (Double
m1, Double
v1, Double
1)
step :: (Double, Double, Double) -> Double -> (Double, Double, Double)
step (Double
m, Double
v, Double
t) Double
g =
let t' :: Double
t' = Double
t Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
1
m' :: Double
m' = Double
beta1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
m Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g
v' :: Double
v' = Double
beta2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
v Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g
in (Double
m', Double
v', Double
t')
extract :: (Double, Double, Double) -> Double
extract (Double
m, Double
v, Double
t) =
let mHat :: Double
mHat = Double
m Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Double
t)
vHat :: Double
vHat = Double
v Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Double
t)
in Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
mHat Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double -> Double
forall a. Floating a => a -> a
sqrt Double
vHat Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
eps)
adamDecomposed :: Double -> Double -> Double -> Double -> [Double] -> [Double]
adamDecomposed :: Double -> Double -> Double -> Double -> [Double] -> [Double]
adamDecomposed Double
alpha Double
beta1 Double
beta2 Double
eps [Double]
gs =
let ms :: [Double]
ms = Int -> [Double] -> [Double]
forall a. Int -> [a] -> [a]
drop Int
1 ([Double] -> [Double]) -> [Double] -> [Double]
forall a b. (a -> b) -> a -> b
$ (Double -> Double -> Double) -> Double -> [Double] -> [Double]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' (\Double
m Double
g -> Double
beta1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
m Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g) Double
0 [Double]
gs
vs :: [Double]
vs = Int -> [Double] -> [Double]
forall a. Int -> [a] -> [a]
drop Int
1 ([Double] -> [Double]) -> [Double] -> [Double]
forall a b. (a -> b) -> a -> b
$ (Double -> Double -> Double) -> Double -> [Double] -> [Double]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' (\Double
v Double
g -> Double
beta2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
v Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g) Double
0 [Double]
gs
update :: Int -> Double -> Double -> Double
update Int
t Double
m Double
v =
let mHat :: Double
mHat = Double
m Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
t)
vHat :: Double
vHat = Double
v Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
t)
in Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
mHat Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double -> Double
forall a. Floating a => a -> a
sqrt Double
vHat Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
eps)
in (Int -> Double -> Double -> Double)
-> [Int] -> [Double] -> [Double] -> [Double]
forall a b c d. (a -> b -> c -> d) -> [a] -> [b] -> [c] -> [d]
zipWith3 Int -> Double -> Double -> Double
update [Int
1 :: Int ..] [Double]
ms [Double]
vs
adamReference :: Double -> Double -> Double -> [Double] -> [(Double, Double, Int)]
adamReference :: Double -> Double -> Double -> [Double] -> [(Double, Double, Int)]
adamReference Double
beta1 Double
beta2 Double
g0 [Double]
gs = Int -> [(Double, Double, Int)] -> [(Double, Double, Int)]
forall a. Int -> [a] -> [a]
drop Int
1 ([(Double, Double, Int)] -> [(Double, Double, Int)])
-> [(Double, Double, Int)] -> [(Double, Double, Int)]
forall a b. (a -> b) -> a -> b
$ ((Double, Double, Int) -> Double -> (Double, Double, Int))
-> (Double, Double, Int) -> [Double] -> [(Double, Double, Int)]
forall b a. (b -> a -> b) -> b -> [a] -> [b]
scanl' (Double, Double, Int) -> Double -> (Double, Double, Int)
step (Double
0, Double
0, Int
0) (Double
g0 Double -> [Double] -> [Double]
forall a. a -> [a] -> [a]
: [Double]
gs)
where
step :: (Double, Double, Int) -> Double -> (Double, Double, Int)
step (Double
m, Double
v, Int
t) Double
g =
let t' :: Int
t' = Int
t Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1
m' :: Double
m' = Double
beta1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
m Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g
v' :: Double
v' = Double
beta2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
v Double -> Double -> Double
forall a. Num a => a -> a -> a
+ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2) Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
g
in (Double
m', Double
v', Int
t')
adamUpdates :: Double -> Double -> Double -> Double -> [(Double, Double, Int)] -> [Double]
adamUpdates :: Double
-> Double
-> Double
-> Double
-> [(Double, Double, Int)]
-> [Double]
adamUpdates Double
alpha Double
beta1 Double
beta2 Double
eps = ((Double, Double, Int) -> Double)
-> [(Double, Double, Int)] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (Double, Double, Int) -> Double
update
where
update :: (Double, Double, Int) -> Double
update (Double
m, Double
v, Int
t) =
let mHat :: Double
mHat = Double
m Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta1 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
t)
vHat :: Double
vHat = Double
v Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double
1 Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
beta2 Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
t)
in Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
mHat Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double -> Double
forall a. Floating a => a -> a
sqrt Double
vHat Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
eps)