src/linear_regression.bend source
src/linear_regression.bend on the hub · documented module
import Baseimport ./batch.bend as Batchimport ./numeric.bend as Numericimport ./stats.bend as Stats# Ordinary least squares with one input feature and a fitted intercept.type Model is Data: Model{slope: F32, intercept: F32}type FitError is Data: NotEnoughSamples{} NonFiniteData{} ConstantFeature{} NumericalFailure{}def error_message(error: FitError) -> String: match error: case NotEnoughSamples{}: "Fit needs at least two samples." case NonFiniteData{}: "Training data contains NaN or infinity." case ConstantFeature{}: "Input feature has zero variance at F32 precision." case NumericalFailure{}: "F32 arithmetic overflowed or became invalid; rescale the data."def checked_model(ok: Bool, slope: F32, intercept: F32) -> Result<&2, &2, FitError, Model>: match ok: case False{}: Fail{NumericalFailure{}} case True{}: Done{Model{slope, intercept}}def solve(nonzero: Bool, mean_x: F32, mean_y: F32, xx: F32, xy: F32) -> Result<&2, &2, FitError, Model>: match nonzero: case False{}: Fail{ConstantFeature{}} case True{}: +slope = F32.div(xy, xx) +intercept = F32.sub(mean_y, F32.mul(slope, mean_x)) checked_model(Numeric.is_finite(slope) && Numeric.is_finite(intercept), slope, intercept)def checked_moments(ok: Bool, mean_x: F32, mean_y: F32, +xx: F32, xy: F32) -> Result<&2, &2, FitError, Model>: match ok: case False{}: Fail{NumericalFailure{}} case True{}: solve(F32.is_gt(xx, 0.0), mean_x, mean_y, xx, xy)def checked_input(valid: Bool, +mean_x: F32, +mean_y: F32, +xx: F32, +xy: F32) -> Result<&2, &2, FitError, Model>: match valid: case False{}: Fail{NonFiniteData{}} case True{}: checked_moments(Numeric.is_finite(mean_x) && Numeric.is_finite(mean_y) && Numeric.is_finite(xx) && Numeric.is_finite(xy), mean_x, mean_y, xx, xy)def finish(stats: Stats.Moments) -> Result<&2, &2, FitError, Model>: match stats: case Stats.NoSamples{}: Fail{NotEnoughSamples{}} case Stats.Moments{n, mean_x, mean_y, xx, xy, valid}: match n: case 0n: Fail{NotEnoughSamples{}} case 1n+0n: Fail{NotEnoughSamples{}} case 1n+1n+rest: checked_input(valid, mean_x, mean_y, xx, xy)def fit(rows: Batch.Batch<Stats.Sample>) -> Result<&2, &2, FitError, Model>: finish(Stats.moments(rows))def predict(model: Model, x: F32) -> F32: match model: case Model{slope, intercept}: F32.add(F32.mul(slope, x), intercept)def predict_batch(+model: Model, rows: Batch.Batch<F32>) -> Batch.Batch<F32>: match rows: case Batch.Empty{}: Batch.Empty{} case Batch.Item{x}: Batch.Item{predict(model, x)} case Batch.Fork{left, right}: a b = predict_batch(model, left) predict_batch(model, right) Batch.Fork{a, b}