-- | Toy 2-layer neural network fit.
--
-- A small regression example: 1→2→1 network with tanh activation, trained
-- by Adam-style gradient descent to approximate @y = 2*x@.
--
-- The oracle in circuits-learn-axioma asserts that loss decreases
-- over training: @loss(weights_0) > loss(weights_N)@.
module Circuit.Learn.Fit
  ( -- * Network
    forward,
    loss,

    -- * Training
    fit,

    -- * Dev dataset
    toyData,
  )
where

import Prelude

-- ---------------------------------------------------------------------------
-- 2-layer network: 1 → 2 → 1 with tanh activation
-- 7 params: w1,w2,b1,b2,w3,w4,b3
-- ---------------------------------------------------------------------------

-- | tanh activation.
tanhAct :: Double -> Double
tanhAct :: Double -> Double
tanhAct Double
x = (Double -> Double
forall a. Floating a => a -> a
exp Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double -> Double
forall a. Floating a => a -> a
exp (-Double
x)) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double -> Double
forall a. Floating a => a -> a
exp Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double -> Double
forall a. Floating a => a -> a
exp (-Double
x))

-- | Forward pass: 1 → 2 → 1.
forward :: [Double] -> Double -> Double
forward :: [Double] -> Double -> Double
forward [Double
w1, Double
w2, Double
b1, Double
b2, Double
w3, Double
w4, Double
b3] Double
x =
  let h1 :: Double
h1 = Double -> Double
tanhAct (Double
w1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b1)
      h2 :: Double
h2 = Double -> Double
tanhAct (Double
w2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b2)
   in Double
w3 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h1 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
w4 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h2 Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
b3
forward [Double]
_ Double
_ = [Char] -> Double
forall a. HasCallStack => [Char] -> a
error [Char]
"forward: expected exactly 7 params"

-- | Mean squared error loss.
loss :: [Double] -> [(Double, Double)] -> Double
loss :: [Double] -> [(Double, Double)] -> Double
loss [Double]
w [(Double, Double)]
dataset =
  let total :: Double
total = [Double] -> Double
forall a. Num a => [a] -> a
forall (t :: * -> *) a. (Foldable t, Num a) => t a -> a
sum [([Double] -> Double -> Double
forward [Double]
w Double
x Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
y) Double -> Double -> Double
forall a. Floating a => a -> a -> a
** Double
2 | (Double
x, Double
y) <- [(Double, Double)]
dataset]
   in Double
total Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Int -> Double
forall a b. (Integral a, Num b) => a -> b
fromIntegral ([(Double, Double)] -> Int
forall a. [a] -> Int
forall (t :: * -> *) a. Foldable t => t a -> Int
length [(Double, Double)]
dataset)

-- | Toy dataset: y ≈ 2*x with slight noise on the last point.
toyData :: [(Double, Double)]
toyData :: [(Double, Double)]
toyData = [(Double
0.0, Double
0.0), (Double
1.0, Double
2.0), (Double
2.0, Double
4.5)]

-- ---------------------------------------------------------------------------
-- Training loop — Adam-style gradient descent
-- ---------------------------------------------------------------------------

-- | Compute the gradient via central finite differences.
grad :: [Double] -> [(Double, Double)] -> [Double]
grad :: [Double] -> [(Double, Double)] -> [Double]
grad [Double]
w [(Double, Double)]
dataset =
  let base :: Double
base = [Double] -> [(Double, Double)] -> Double
loss [Double]
w [(Double, Double)]
dataset
      h :: Double
h = Double
1e-5
      perturb :: Int -> Double -> Double
perturb Int
i Double
sign =
        let w' :: [Double]
w' = (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith Double -> Double -> Double
forall a. Num a => a -> a -> a
(+) [Double]
w (Int -> Double -> [Double]
forall a. Int -> a -> [a]
replicate Int
i Double
0 [Double] -> [Double] -> [Double]
forall a. [a] -> [a] -> [a]
++ [Double
sign Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
h] [Double] -> [Double] -> [Double]
forall a. [a] -> [a] -> [a]
++ Double -> [Double]
forall a. a -> [a]
repeat Double
0)
         in Double
sign Double -> Double -> Double
forall a. Num a => a -> a -> a
* ([Double] -> [(Double, Double)] -> Double
loss [Double]
w' [(Double, Double)]
dataset Double -> Double -> Double
forall a. Num a => a -> a -> a
- Double
base) Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ Double
h
   in [Int -> Double -> Double
perturb Int
i Double
1 | Int
i <- [Int
0 .. Int
6]]

-- | Adam step on a (params, m, v, t) state tuple.
adamStep ::
  Double ->
  Double ->
  Double ->
  [(Double, Double)] ->
  ([Double], [Double], [Double], Int) ->
  ([Double], [Double], [Double], Int)
adamStep :: Double
-> Double
-> Double
-> [(Double, Double)]
-> ([Double], [Double], [Double], Int)
-> ([Double], [Double], [Double], Int)
adamStep Double
beta1 Double
beta2 Double
alpha [(Double, Double)]
dataset ([Double]
w, [Double]
m, [Double]
v, Int
t) =
  let g :: [Double]
g = [Double] -> [(Double, Double)] -> [Double]
grad [Double]
w [(Double, Double)]
dataset
      t' :: Int
t' = Int
t Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
1
      m' :: [Double]
m' = (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (\Double
mi Double
gi -> Double
beta1 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
mi 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
gi) [Double]
m [Double]
g
      v' :: [Double]
v' = (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (\Double
vi Double
gi -> Double
beta2 Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
vi 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
gi Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
gi) [Double]
v [Double]
g
      mHat :: [Double]
mHat = (Double -> Double) -> [Double] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (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')) [Double]
m'
      vHat :: [Double]
vHat = (Double -> Double) -> [Double] -> [Double]
forall a b. (a -> b) -> [a] -> [b]
map (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')) [Double]
v'
      update :: [Double]
update = (Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (\Double
mh Double
vh -> Double
alpha Double -> Double -> Double
forall a. Num a => a -> a -> a
* Double
mh Double -> Double -> Double
forall a. Fractional a => a -> a -> a
/ (Double -> Double
forall a. Floating a => a -> a
sqrt Double
vh Double -> Double -> Double
forall a. Num a => a -> a -> a
+ Double
eps)) [Double]
mHat [Double]
vHat
      eps :: Double
eps = Double
1e-8
   in ((Double -> Double -> Double) -> [Double] -> [Double] -> [Double]
forall a b c. (a -> b -> c) -> [a] -> [b] -> [c]
zipWith (-) [Double]
w [Double]
update, [Double]
m', [Double]
v', Int
t')

-- | Run N steps of Adam, returning trained params and initial/final loss.
fit ::
  Int ->
  [Double] ->
  [(Double, Double)] ->
  ([Double], Double, Double)
fit :: Int -> [Double] -> [(Double, Double)] -> ([Double], Double, Double)
fit Int
steps [Double]
w0 [(Double, Double)]
dataset =
  let beta1 :: Double
beta1 = Double
0.9
      beta2 :: Double
beta2 = Double
0.999
      alpha :: Double
alpha = Double
0.05
      initLoss :: Double
initLoss = [Double] -> [(Double, Double)] -> Double
loss [Double]
w0 [(Double, Double)]
dataset
      initState :: ([Double], [Double], [Double], Int)
initState = ([Double]
w0, Int -> Double -> [Double]
forall a. Int -> a -> [a]
replicate Int
7 Double
0, Int -> Double -> [Double]
forall a. Int -> a -> [a]
replicate Int
7 Double
0, Int
0 :: Int)
      ([Double]
wN, [Double]
_, [Double]
_, Int
_) = (([Double], [Double], [Double], Int)
 -> ([Double], [Double], [Double], Int))
-> ([Double], [Double], [Double], Int)
-> [([Double], [Double], [Double], Int)]
forall a. (a -> a) -> a -> [a]
iterate (Double
-> Double
-> Double
-> [(Double, Double)]
-> ([Double], [Double], [Double], Int)
-> ([Double], [Double], [Double], Int)
adamStep Double
beta1 Double
beta2 Double
alpha [(Double, Double)]
dataset) ([Double], [Double], [Double], Int)
initState [([Double], [Double], [Double], Int)]
-> Int -> ([Double], [Double], [Double], Int)
forall a. HasCallStack => [a] -> Int -> a
!! Int
steps
      finalLoss :: Double
finalLoss = [Double] -> [(Double, Double)] -> Double
loss [Double]
wN [(Double, Double)]
dataset
   in ([Double]
wN, Double
initLoss, Double
finalLoss)