0 | module NN.Optimisers.Instances
  1 |
  2 | import Data.Materialise
  3 | import Data.Container.Additive
  4 | import Data.Num
  5 | import NN.Optimisers.Definition
  6 | import NN.Utils
  7 |
  8 | ||| Gradient descent optimiser. Has trivial state
  9 | ||| @lr is the learning rate
 10 | public export
 11 | GD : Neg pType =>
 12 |   (mon : ComMonoid pType) => FromDouble pType =>
 13 |   {default 0.001 lr : pType} -> Optimiser (Const pType) Unit
 14 | GD = MkOptimiser
 15 |   (!% \(p, ()) => (p ** \p' => (p - lr * p', ())))
 16 |   (pure ())
 17 |
 18 | ||| Gradient ascent optimiser. Has trivial state
 19 | ||| @lr is the learning rate
 20 | public export
 21 | GA : Neg pType =>
 22 |   (mon : ComMonoid pType) => FromDouble pType =>
 23 |   {default 0.001 lr : pType} -> Optimiser (Const pType) Unit
 24 | GA = MkOptimiser
 25 |   (!% \(p, ()) => (p ** \p' => (p + lr * p', ())))
 26 |   (pure ())
 27 |
 28 | namespace Momentum
 29 |   public export
 30 |   momentumUpdate : Neg pType =>
 31 |     {lr : pType} ->
 32 |     (gamma : pType) ->
 33 |     (p : pType) ->
 34 |     (s : pType) ->
 35 |     (p' : pType) ->
 36 |     (pType, pType)
 37 |   momentumUpdate gamma p s p' = let s' = gamma * s + p'
 38 |                                 in (p - lr * s', s')
 39 |
 40 |   public export
 41 |   lookAhead : Num pType =>
 42 |     (gamma, p, s : pType) ->
 43 |     pType
 44 |   lookAhead gamma p s = p + gamma * s
 45 |   
 46 |   ||| Gradient Descent with momentum, optionally with Nesterov acceleration
 47 |   public export
 48 |   GDMomentum : Neg pType =>
 49 |    (mon : ComMonoid pType) =>
 50 |    FromDouble 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))
 58 |     (pure 0)
 59 |   
 60 | namespace Adam
 61 |   ||| Adam step. The moments are parameter-shaped; the bias-correction powers
 62 |   ||| beta^t are the same scalar at every coordinate, so they are `Double`
 63 |   public export
 64 |   adamUpdate : Neg pType => Fractional pType => Sqrt pType =>
 65 |     FromDouble pType => Materialise pType =>
 66 |     {lr : pType} ->
 67 |     (beta1 : Double) ->
 68 |     (beta2 : Double) ->
 69 |     (epsilon : pType) ->
 70 |     (p : pType) ->
 71 |     (m : pType) ->
 72 |     (v : pType) ->
 73 |     (b1p : Double) ->
 74 |     (b2p : Double) ->
 75 |     (g : 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')
 81 |         b1p' = b1p * beta1
 82 |         b2p' = b2p * beta2
 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')
 86 |
 87 |   ||| Adam optimiser (Kingma & Ba, 2014)
 88 |   ||| State: the two moments and the two scalar bias-correction powers
 89 |   ||| @lr is the learning rate
 90 |   ||| @beta1 is the exponential decay rate for the first moment estimate
 91 |   ||| @beta2 is the exponential decay rate for the second moment estimate
 92 |   ||| @epsilon is a small constant for numerical stability
 93 |   public export
 94 |   Adam : Neg pType =>
 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)
103 |   Adam = MkOptimiser
104 |     (!% \(p, (m, v, b1p, b2p)) =>
105 |       (p ** adamUpdate {lr} beta1 beta2 epsilon p m v b1p b2p))
106 |     (pure (0, 0, 1, 1))