0 | module NN.Training.Examples.LinearRegression
5 | 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 : IsFlat (GetParam f)) =>
31 | linearRegression f@(MkPara (MkAddCont (Const p)) _)
32 | {isFlat = MkIsFlat p @{mon}} numSteps = do
33 | putStrLn "Training a linear regression model..."
34 | trainData <- linearRegressionDataLoader
35 | testDataLoader <- makeDataLoader [20, 50, 100] (pure . groundTruth)
36 | pTrained <- fst <$> optimise
37 | {l=Const Double, e=pushDown (Const Double >< Const Double)}
38 | (buildSupervisedLearningSystem f SquaredDifference)
39 | (handleData trainData)
40 | (GDMomentum {pType=(GetParam f).Shp})
42 | fromCostate (eval f pTrained) (snd $
inputs testDataLoader)
43 | avgLoss <- fromCostate (averageLoss f SquaredDifference pTrained) (dataset testDataLoader)
44 | putStrLn "Average loss: \{showSci avgLoss}"