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 | initParam : IO paramCont.Shp
36 | initState : IO stateTy
39 | (.fwd) : Optimiser p s -> (p.Shp, s) -> p.Shp
40 | (.fwd) (MkOptimiser opt _ _) = opt.fwd
43 | (.bwd) : (opt : Optimiser pCont stateTy) ->
44 | (ps : (pCont.Shp, stateTy)) ->
45 | (pCont.Pos (opt.fwd ps)) ->
46 | (pCont.Shp, stateTy)
47 | (.bwd) (MkOptimiser opt _ _) = opt.bwd
52 | composeParallel : Optimiser pCont s ->
53 | Optimiser qCont t ->
54 | Optimiser (pCont >< qCont) (s, t)
55 | composeParallel (MkOptimiser o1 initP initS) (MkOptimiser o2 initQ initT) = MkOptimiser
56 | (!% \((p, q), (s, t)) => (
(o1.fwd (p, s), o2.fwd (q, t)) **
57 | \(p', q') => let (pUpdated, sUpdated) = o1.bwd (p, s) p'
58 | (qUpdated, tUpdated) = o2.bwd (q, t) q'
59 | in ((pUpdated, qUpdated), (sUpdated, tUpdated)))
)
60 | [| (initP, initQ) |]
61 | [| (initS, initT) |]