///|
// Learner injection for v0.54.0+.
//
// Upstream `doubleml-for-py` accepts any sklearn-compatible
// estimator as `learner_l` / `learner_m`. The MoonBit port
// already declares the `Learner` trait in `linear.mbt` with
// two methods (`fit` + `predict`) and provides a generic
// `cross_fit_predict[T : Learner]` cross-fit helper. This
// module adds two learner variants on top of the existing
// `LinearRegression` (which is the v0.54.0 default):
//
// - `ConstantLearner(value)`: predicts `value` for every
// row. Useful for unit tests that want a known-null
// nuisance estimator (the DML score should collapse to
// θ ≈ 0 when `l_hat == 0`).
//
// - `NoopLearner`: a trivial learner that always returns
// zero for every row. Useful for sanity tests that the
// DML pipeline doesn't divide by zero when the nuisance
// is trivial.
//
// Adding more learner types (logistic regression, gradient
// boosting, random forest) is a one-line `impl Learner for X`
// away — no other code changes required, because
// `cross_fit_predict[T : Learner]` is generic.
//
// `LearnerDispatch` is the enum wrapper that lets a non-generic
// struct field hold one of the concrete learner types. The
// dispatcher `cross_fit_predict_dispatch` pattern-matches the
// enum and routes each arm to the generic
// `cross_fit_predict[T : Learner]` with the right concrete
// `T`. Adding a new learner type means:
// 1. declare `impl Learner for NewType`
// 2. add a `NewType(NewType)` arm to `LearnerDispatch`
// 3. add the corresponding match arm in
// `cross_fit_predict_dispatch` and
// `learner_dispatch_default`.
///|
/// Constant-predictor learner. `fit` is a no-op (returns
/// `self`); `predict` returns a length-`x.rows()` array of
/// the constant `value` for every row, ignoring `x`.
pub struct ConstantLearner {
value : Double
} derive(Debug)
///|
pub extend ConstantLearner with @moonbitlang/core/debug.Debug::{to_repr}
///|
pub extend ConstantLearner with Learner::{fit, predict}
///|
pub fn ConstantLearner::new(value : Double) -> ConstantLearner {
{ value, }
}
///|
impl Learner for ConstantLearner with fn fit(self, x, y) {
// "fit" for a constant predictor: stash the n_rows + n_targets
// shape on `self` so subsequent `predict` calls can short-circuit
// (no actual training occurs; the predictor is constant).
// For now, we just consume the inputs and return self.
let _n : Int = x.rows() * y.length()
let _ = _n
self
}
///|
impl Learner for ConstantLearner with fn predict(self, x) {
Array::make(x.rows(), self.value)
}
///|
/// Trivial learner: predicts zero for every row. `fit` is a
/// no-op; `predict` returns all zeros. Useful for sanity
/// tests that the DML pipeline doesn't divide by zero when
/// the nuisance is trivial.
pub struct NoopLearner {
// `value : Double = 0.0` placeholder; preserved as a struct
// (not a unit type) so adding fields later is non-breaking.
value : Double
} derive(Debug)
///|
pub extend NoopLearner with @moonbitlang/core/debug.Debug::{to_repr}
///|
pub extend NoopLearner with Learner::{fit, predict}
///|
pub fn NoopLearner::new() -> NoopLearner {
{ value: 0.0, }
}
///|
impl Learner for NoopLearner with fn fit(self, x, y) {
// Same as ConstantLearner::fit: touch the inputs so the
// compiler doesn't flag the impl body as unused, then
// return self.
let _n : Int = x.rows() * y.length()
let _ = _n
self
}
///|
impl Learner for NoopLearner with fn predict(self, x) {
Array::make(x.rows(), self.value)
}
///|
/// Dispatch enum for the `Learner` trait. Stores one of the
/// concrete learner types and lets `cross_fit_predict_dispatch`
/// route to the right `cross_fit_predict[T : Learner]` arm.
///
/// The enum keeps `DoubleMLPLR` non-generic in its struct
/// definition (callers don't have to write
/// `DoubleMLPLR[LinearRegression, LinearRegression]`), at the
/// cost of an extra match arm per learner type. Adding a new
/// learner is 3 small edits (see the file header).
pub enum LearnerDispatch {
/// OLS via `LinearRegression` (the v0.54.0 default)
LinearRegression(LinearRegression)
/// Constant-predictor learner (see `ConstantLearner`)
Constant(ConstantLearner)
/// Zero-predictor learner (see `NoopLearner`)
Noop(NoopLearner)
/// Pure-MoonBit random forest (Breiman 2001 regression).
/// v0.56.0+: Path A — first non-OLS learner.
RandomForest(RFLearner)
/// Pure-MoonBit gradient boosting (Friedman 2001 regression).
/// v0.57.0+: Path A — second non-OLS learner.
GradientBoosting(GBLearner)
} derive(Debug)
///|
pub extend LearnerDispatch with @moonbitlang/core/debug.Debug::{to_repr}
///|
/// Default `LearnerDispatch` (a fresh `LinearRegression`).
/// Used as the default value for the `learner_l` /
/// `learner_m` labeled params on `DoubleMLPLR::new`.
pub fn LearnerDispatch::linear_regression() -> LearnerDispatch {
LinearRegression(LinearRegression::new())
}
///|
/// Wrap a `ConstantLearner` in a `LearnerDispatch`.
pub fn LearnerDispatch::constant(value : Double) -> LearnerDispatch {
Constant(ConstantLearner::new(value))
}
///|
/// Wrap a `NoopLearner` in a `LearnerDispatch`.
pub fn LearnerDispatch::noop() -> LearnerDispatch {
Noop(NoopLearner::new())
}
///|
/// Wrap an `RFLearner` in a `LearnerDispatch`. v0.56.0+.
/// Factory for the Path A random-forest learner; callers can
/// override n_trees / max_depth / min_samples_leaf / mtry via
/// `RFLearner::new(...)` before wrapping.
pub fn LearnerDispatch::random_forest(rf : RFLearner) -> LearnerDispatch {
RandomForest(rf)
}
///|
/// Wrap a `GBLearner` in a `LearnerDispatch`. v0.57.0+.
/// Factory for the Path A gradient-boosting learner; callers
/// can override n_trees / learning_rate / max_depth / subsample
/// via `GBLearner::new(...)` before wrapping.
pub fn LearnerDispatch::gradient_boosting(gb : GBLearner) -> LearnerDispatch {
GradientBoosting(gb)
}
///|
/// Cross-fit a `LearnerDispatch` over the given folds. This
/// is the dispatch wrapper around `cross_fit_predict[T : Learner]`
/// (which is generic in `T`); each match arm binds a concrete
/// `T` so the generic call resolves correctly.
///
/// Adding a new learner type: add the corresponding match arm
/// here (and the enum variant in `LearnerDispatch`).
pub fn cross_fit_predict_dispatch(
learner : LearnerDispatch,
x : Matrix,
y : Array[Double],
folds : Array[Fold],
) -> Array[Double] {
match learner {
LinearRegression(lr) => cross_fit_predict(lr, x, y, folds)
Constant(c) => cross_fit_predict(c, x, y, folds)
Noop(n) => cross_fit_predict(n, x, y, folds)
RandomForest(rf) => cross_fit_predict(rf, x, y, folds)
GradientBoosting(gb) => cross_fit_predict(gb, x, y, folds)
}
}
///|
/// Single-fit + single-predict helper for a `LearnerDispatch`.
/// Used by estimators (e.g. `DoubleMLSSM` for the augmented
/// `(X, D) -> S` pi fit, `DoubleMLIIVM` for conditional subsets
/// where the augmenting column changes the feature matrix shape)
/// that need a one-off `fit` on a subset followed by a `predict`
/// on a different-shape test set — i.e. cases where
/// `cross_fit_predict_dispatch` (which slices a single matrix
/// by train/test indices) doesn't apply.
///
/// v0.59.0+.
pub fn fit_predict_one_dispatch(
learner : LearnerDispatch,
x_train : Matrix,
y_train : Array[Double],
x_test : Matrix,
) -> Array[Double] {
match learner {
LinearRegression(lr) => lr.fit(x_train, y_train).predict(x_test)
Constant(c) => c.fit(x_train, y_train).predict(x_test)
Noop(n) => n.fit(x_train, y_train).predict(x_test)
RandomForest(rf) => rf.fit(x_train, y_train).predict(x_test)
GradientBoosting(gb) => gb.fit(x_train, y_train).predict(x_test)
}
}