0 | module NN.Optimisers.Instances
2 | import Data.Materialise
3 | import Data.Container.Additive
5 | import NN.Optimisers.Definition
12 | (mon : ComMonoid pType) => FromDouble pType =>
13 | {default 0.001 lr : pType} -> Optimiser (Const pType) Unit
15 | (!% \(p, ()) => (
p ** \p' => (p - lr * p', ()))
)
22 | (mon : ComMonoid pType) => FromDouble pType =>
23 | {default 0.001 lr : pType} -> Optimiser (Const pType) Unit
25 | (!% \(p, ()) => (
p ** \p' => (p + lr * p', ()))
)
30 | momentumUpdate : Neg pType =>
37 | momentumUpdate gamma p s p' = let s' = gamma * s + p'
38 | in (p - lr * s', s')
41 | lookAhead : Num pType =>
42 | (gamma, p, s : pType) ->
44 | lookAhead gamma p s = p + gamma * s
48 | GDMomentum : Neg pType =>
49 | (mon : ComMonoid pType) =>
51 | {default False nesterov : Bool} ->
52 | {default 0.001 lr : pType} ->
53 | {default 0.9 gamma : pType} ->
54 | Optimiser (Const pType) pType
55 | GDMomentum = MkOptimiser
56 | (!% \(p, s) => (
if nesterov then lookAhead gamma p s else p
57 | ** momentumUpdate {lr} gamma p s)
)
64 | adamUpdate : Neg pType => Fractional pType => Sqrt pType =>
65 | FromDouble pType => Materialise pType =>
69 | (epsilon : pType) ->
76 | (pType, pType, pType, Double, Double)
77 | adamUpdate beta1 beta2 epsilon p m v b1p b2p g =
78 | let g' = materialise g
79 | m' = materialise (fromDouble beta1 * m + fromDouble (1 - beta1) * g')
80 | v' = materialise (fromDouble beta2 * v + fromDouble (1 - beta2) * g' * g')
83 | mHat = fromDouble (1 / (1 - b1p')) * m'
84 | vHat = fromDouble (1 / (1 - b2p')) * v'
85 | in (p - lr * mHat / (sqrt vHat + epsilon), m', v', b1p', b2p')
95 | (mon : ComMonoid pType) =>
96 | FromDouble pType => Materialise pType =>
97 | Fractional pType => Sqrt pType =>
98 | {default 0.001 lr : pType} ->
99 | {default 0.9 beta1 : Double} ->
100 | {default 0.999 beta2 : Double} ->
101 | {default 1.0e-8 epsilon : pType} ->
102 | Optimiser (Const pType) (pType, pType, Double, Double)
104 | (!% \(p, (m, v, b1p, b2p)) =>
105 | (
p ** adamUpdate {lr} beta1 beta2 epsilon p m v b1p b2p)
)
106 | (pure (0, 0, 1, 1))