0 | module NN.Training.Training
3 | import Data.Container.Additive as Additive
4 | import public Data.ScientificNotation
8 | import NN.Training.DataLoader
10 | import Data.Autodiff.Model
41 | optimiseStep : {p, l : AddCont} -> {e : Cont} ->
42 | InterfaceOnPositions l Num =>
43 | (f : p =%+> e >-+@ l) ->
44 | (handleEffect : Costate (IO <!> e)) ->
45 | (optimiser : Optimiser p stateTy) ->
46 | Costate (IO <!> (Const (p.Shp, stateTy)))
47 | optimiseStep f handleEffect (MkOptimiser opt _) =
48 | let closeFunction : p =%+> !* e
49 | closeFunction = f %+>> (id >-+@ constantOne) %+>> actionToFree
51 | closeFunctionT : UC p =%> e
52 | closeFunctionT = addContTransposeInv closeFunction
54 | in (IO <!> (opt %>> closeFunctionT)) %>> handleEffect
58 | evalFw : {0 e : Cont} ->
59 | (f : a -> Ext e b) ->
60 | (handleEffect : Costate (IO <!> e)) ->
61 | Costate (IO <!> (Const2 a b))
62 | evalFw f handleEffect = toCostate $
\ps => do
63 | let (eInp <| outGivenEffect) = f ps
64 | e <- fromCostate handleEffect eInp
65 | pure $
outGivenEffect e
70 | optimise : {p, l : AddCont} -> {e : Cont} ->
71 | InterfaceOnPositions l Num =>
72 | {default 100 printEvery : Nat} ->
73 | Materialise p.Shp => Materialise stateTy =>
74 | ScientificDisplay p.Shp => ScientificDisplay l.Shp => ScientificDisplay stateTy =>
75 | (f : p =%+> e >-+@ l) ->
76 | (handleEffect : Costate (IO <!> e)) ->
77 | (initParam : IO p.Shp) ->
78 | (opt : Optimiser p stateTy) ->
81 | optimise f handleEffect initParam opt numSteps = do
82 | currentValue <- initParam
83 | currentState <- opt.initState
84 | runActionUntilMaxSteps
86 | {printEvery=printEvery}
87 | (fromCostate $
optimiseStep f handleEffect opt)
90 | (currentValue, currentState)
91 | (fromCostate $
evalFw (f.fwd . opt.fwd) handleEffect)
95 | buildSupervisedLearningSystem : (f : x =\\=> y) -> (loss : y =\\=> l) ->
96 | Materialise (Param f).Shp => InterfaceOnPositions (Param f) Materialise =>
97 | Param f =%+> (SupervisedData x.Shp (Param loss).Shp) >-+@ l
98 | buildSupervisedLearningSystem f loss =
99 | let supplied : (a >*< b) >*< c =%+> a >*< (c >*< b)
100 | supplied = assocL %+>> (id >*< swap)
101 | in materialiseCont %+>> pushIntoContinuation {d=x>*<Param loss}
102 | (supplied %+>> (composePara f loss).Run)
105 | namespace WithEffect
110 | totalLoss : Num l.Shp =>
111 | (f : x =\\=> e >-+@ y) ->
112 | (loss : y =\\=> l) ->
113 | (p : (Param f).Shp) ->
114 | (handleEffect : Costate (IO <!> e)) ->
115 | Costate (IO <!> (Const2 (Vect n (x.Shp, (Param loss).Shp)) l.Shp))
116 | totalLoss (MkPara pCont f) (MkPara z loss) p handleEffect
117 | = let evalFWithLoss : (x.Shp, z.Shp) -> IO l.Shp
118 | evalFWithLoss (x, yTrue) = do
119 | yPred <- fromCostate (evalFw f.fwd handleEffect) (x, p)
120 | pure $
loss.fwd (yPred, yTrue)
122 | in toCostate $
\testData => do
123 | losses <- traverse evalFWithLoss testData
124 | pure $
Prelude.sum losses
128 | averageLoss : {n : Nat} ->
129 | Num l.Shp => Fractional l.Shp => Cast Nat l.Shp =>
130 | (f : x =\\=> e >-+@ y) ->
131 | (loss : y =\\=> l) ->
132 | (p : (Param f).Shp) ->
133 | (handleEffect : Costate (IO <!> e)) ->
134 | Costate (IO <!> (Const2 (Vect n (x.Shp, (Param loss).Shp)) l.Shp))
135 | averageLoss f loss p handleEffect = toCostate $
\testData => do
136 | lossSum <- fromCostate (totalLoss f loss p handleEffect) testData
137 | pure (lossSum / cast n)
141 | train : {a, b : AddCont} -> {l : AddCont} ->
142 | {default 100 printEvery : Nat} ->
144 | (loss : b =\\=> l) ->
145 | InterfaceOnPositions l Num =>
146 | ScientificDisplay l.Shp =>
147 | Materialise m.Params => Materialise stateTy =>
148 | ScientificDisplay m.Params => ScientificDisplay stateTy =>
149 | (trainData : DataLoader a.Shp (Param loss).Shp) ->
150 | (opt : Optimiser (ParamCont m) stateTy) ->
151 | (numSteps : Nat) ->
152 | IO (m.Params, stateTy)
153 | train m loss trainData opt numSteps = optimise {printEvery}
154 | (buildSupervisedLearningSystem (toPara m) loss)
155 | (handleData trainData)
162 | averageLoss : {0 a, b, l : AddCont} ->
164 | (loss : b =\\=> l) ->
165 | Fractional l.Shp => Cast Nat l.Shp =>
167 | (dl : DataLoader a.Shp (Param loss).Shp) ->
169 | averageLoss m (MkPara z loss) p dl =
170 | let pointLoss : (a.Shp, z.Shp) -> l.Shp
171 | pointLoss (x, yTrue) = loss.fwd (m.fwd x p, yTrue)
172 | in Prelude.sum (pointLoss <$> dl.dataset) / cast dl.datasetSize
176 | evalPrint : {0 a, b : AddCont} ->
177 | ScientificDisplay a.Shp => ScientificDisplay b.Shp =>
178 | (m : a -\-> b) -> (p : m.Params) ->
179 | DataLoader a.Shp b.Shp -> IO ()
180 | evalPrint m p dl = for_ dl.dataset $
\(x, _) =>
181 | putStrLn "Input: \{showSci x}, Predicted: \{showSci (m.fwd x p)}"