module Circuit.Learn.Fit
(
forward,
loss,
fit,
toyData,
)
where
import Prelude
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 :: [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"
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)
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)]
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]]
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')
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)