0 | module NN.Optimisers.Definition
3 | import Data.Container.Additive
15 | record DOptimiser {inputTy : Type}
16 | (paramCont : inputTy -> AddCont)
19 | constructor MkDOptimiser
21 | opt : (x : inputTy) ->
22 | (Const (paramCont x).Shp >< Const stateTy) =%> UC (paramCont x)
28 | (paramCont : AddCont)
31 | constructor MkOptimiser
32 | opt : Const paramCont.Shp >< Const stateTy =%> UC paramCont
34 | initState : IO stateTy
38 | (.fwd) : Optimiser p s -> (p.Shp, s) -> p.Shp
39 | (.fwd) (MkOptimiser opt _) = opt.fwd
42 | (.bwd) : (opt : Optimiser pCont stateTy) ->
43 | (ps : (pCont.Shp, stateTy)) ->
44 | (pCont.PosSet (opt.fwd ps)) ->
45 | (pCont.Shp, stateTy)
46 | (.bwd) (MkOptimiser opt _) = opt.bwd
51 | composeParallel : Optimiser pCont s ->
52 | Optimiser qCont t ->
53 | Optimiser (pCont >*< qCont) (s, t)
54 | composeParallel (MkOptimiser o1 initS) (MkOptimiser o2 initT) = MkOptimiser
55 | (!% \((p, q), (s, t)) => (
(o1.fwd (p, s), o2.fwd (q, t)) **
56 | \(p', q') => let (pUpdated, sUpdated) = o1.bwd (p, s) p'
57 | (qUpdated, tUpdated) = o2.bwd (q, t) q'
58 | in ((pUpdated, qUpdated), (sUpdated, tUpdated)))
)
59 | [| (initS, initT) |]