module Circuit.Learn.Adam
  ( -- * EWMA building block
    ewma,
    ewmaDirect,

    -- * Adam
    adam,
    adamDecomposed,
    adamReference,
    adamUpdates,
  )
where

import Circuit.Process (Process (..), scan)
import Data.List (scanl')

-- | Exponentially weighted moving average as a 'Process'.
--
-- State is the current EWMA value; output is the same value.
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

-- | Direct EWMA recurrence on a list (no 'Process' overhead).
-- First element is after first observation.
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 parameter update as a 'Process'.
--
-- Input: gradient. Output: parameter update (not the new parameter).
-- State: @(m, v, t)@ — first moment, second moment, timestep.
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)

-- | Adam decomposed into two independent EWMA channels (the frontier claim).
--
-- The m channel is an EWMA of raw gradients (@beta1@-weighted), the v channel
-- is an EWMA of squared gradients (@beta2@-weighted).  Each is computed
-- independently via 'scanl'', then combined with the bias-corrected quotient.
-- This produces exactly the same updates as the monolithic 'adam' Process.
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

-- | Reference Adam state recurrence on a gradient trace.
--
-- Forward recurrence starting from @(0, 0, 0)@; the first element corresponds
-- to the state after observing @g0@.
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')

-- | Convert reference states to Adam parameter updates.
--
-- This applies the same bias-correction and scaling as 'adam' so the two can
-- be compared directly on a fixed trace.
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)