0 | module NN.Training.Examples.LinearRegression
4 | import NN.Architectures
9 | exampleInputs : Vect 5 Double
10 | exampleInputs = [1, 2, 3, 4, 5]
13 | groundTruth : Double -> Double
14 | groundTruth x = 2 * x + 1
17 | linearRegressionDataLoader : Monad m => m (DataLoader Double Double)
18 | linearRegressionDataLoader = makeDataLoader exampleInputs (pure . groundTruth)
21 | linearRegression : (m : Const Double -\-> Const Double) ->
22 | Neg m.Params => FromDouble m.Params => ScientificDisplay m.Params =>
23 | Materialise m.Params =>
25 | {default 1000 printEvery : Nat} ->
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}
37 | evalPrint m pTrained testDataLoader
38 | let avgLoss = Model.averageLoss m SquaredError pTrained testDataLoader
39 | putStrLn "Average loss: \{showSci avgLoss}"