0 | module NN.Training.Examples.LinearRegression
6 | import NN.Architectures
11 | exampleInputs : Vect 5 Double
12 | exampleInputs = [1, 2, 3, 4, 5]
15 | groundTruth : Double -> Double
16 | groundTruth x = 2 * x + 1
19 | linearRegressionDataLoader : Monad m => m (DataLoader Double Double)
20 | linearRegressionDataLoader = makeDataLoader exampleInputs (pure . groundTruth)
23 | linearRegression : (f : ParaAddDLens (Const Double) (Const Double)) ->
24 | Neg (GetParam f).Shp => Fractional (GetParam f).Shp =>
25 | Sqrt (GetParam f).Shp =>
26 | Random (GetParam f).Shp =>
27 | FromDouble (GetParam f).Shp => ScientificDisplay (GetParam f).Shp =>
28 | (isFlat : IsConst (GetParam f)) =>
30 | {default 1000 printEvery : Nat} ->
32 | linearRegression f@(MkPara (MkAddCont (Const p)) _)
33 | {isFlat = MkIsConst p @{mon}} numSteps = do
34 | putStrLn "Training a linear regression model..."
35 | trainData <- linearRegressionDataLoader
36 | testDataLoader <- makeDataLoader [20, 50, 100] (pure . groundTruth)
37 | pTrained <- fst <$> optimise
38 | {printEvery=printEvery}
39 | {l=Const Double, e=SupervisedData Double Double}
40 | (buildSupervisedLearningSystem f SquaredDifference)
41 | (handleData trainData)
42 | (GDMomentum {pType=(GetParam f).Shp})
44 | fromCostate (eval f pTrained) (snd $
inputs testDataLoader)
45 | avgLoss <- fromCostate (averageLoss f SquaredDifference pTrained) (dataset testDataLoader)
46 | putStrLn "Average loss: \{showSci avgLoss}"