0 | module NN.Training.Examples.LinearRegression
 1 |
 2 | import Data.Tensor
 3 | import Data.Autodiff
 4 | import NN.Architectures
 5 | import NN.Optimisers
 6 | import NN.Training
 7 |
 8 | public export
 9 | exampleInputs : Vect 5 Double
10 | exampleInputs = [1, 2, 3, 4, 5]
11 |
12 | public export
13 | groundTruth : Double -> Double
14 | groundTruth x = 2 * x + 1
15 |
16 | public export
17 | linearRegressionDataLoader : Monad m => m (DataLoader Double Double)
18 | linearRegressionDataLoader = makeDataLoader exampleInputs (pure . groundTruth)
19 |
20 | public export
21 | linearRegression : (m : Const Double -\-> Const Double) ->
22 |   Neg m.Params => FromDouble m.Params => ScientificDisplay m.Params =>
23 |   Materialise m.Params =>
24 |   (numSteps : Nat) ->
25 |   {default 1000 printEvery : Nat} ->
26 |   IO Double
27 | linearRegression m@(MkModel p @{mon} _ _) numSteps = do
28 |   putStrLn "Training a linear regression model..."
29 |   trainData <- linearRegressionDataLoader
30 |   testDataLoader <- makeDataLoader [20, 50, 100] (pure . groundTruth)
31 |   pTrained <- fst <$> train {printEvery}
32 |     m
33 |     SquaredError
34 |     trainData
35 |     GDMomentum
36 |     numSteps
37 |   evalPrint m pTrained testDataLoader
38 |   let avgLoss = Model.averageLoss m SquaredError pTrained testDataLoader
39 |   putStrLn "Average loss: \{showSci avgLoss}"
40 |   pure avgLoss