///|
/// `TuneParam` — one row of the candidate grid passed to
/// `DoubleML*::tune`. Holds a concrete nuisance-learner
/// combination (one first-slot learner, one second-slot learner)
/// to be scored against the others in the grid.
///
/// The grid is `Array[TuneParam]` (typed, not `Dict`). Reasons
/// for the typed wrapper vs a `Dict[String, LearnerDispatch]`:
///
/// - compile-time key/value type safety (Dict keys are
/// runtime-checked strings)
/// - no Dict-construction FFI overhead per candidate
/// - cleaner documentation (`TuneParam::{learner_l, learner_m}`
/// is a documented struct, not a magic-key dict)
///
/// Python upstream `doubleml.DoubleMLPLR.tune` uses a similar
/// shape internally (a list of `{"learner_l": ..., "learner_m": ...}`
/// type) but exposes it via kwargs to `set_tune_params`.
///
/// # Slot semantics - read this before filling a grid
///
/// The fields are named for the **partially linear outcome** family.
/// Several estimator families reuse the same two slots with a
/// different meaning, and the field names do **not** change with the
/// family:
///
/// | family | first slot (`learner_l`) actually holds | second slot (`learner_m`) |
/// |---|---|---|
/// | `PLR` `PLIV` `LPLR` `PLPR` | outcome learner `E[y\|x]` | nuisance `E[m\|x]` |
/// | `APO` `APOS` `IRM` `IIVM` `CVAR` `PQ` `QTE` `SSM` `DID*` | treatment / propensity learner `E[g\|d,x]` | outcome learner `E[m0+d·m1\|x]` |
/// | `RDD` | discontinuity-side learner `E[y\|x]` | unused (RDD has a single slot) |
///
/// v0.108.0: the pre-existing doc comments in `apo.mbt` / `irm.mbt` /
/// `iivm.mbt` described their pair as `(learner_g, learner_m)`, which
/// reads as a promise that the field is named `learner_g`. It is not
/// - the field is `learner_l` and always has been. Those comments now
/// name the slot's *role* instead. Use `TuneParam::for_outcome` /
/// `TuneParam::for_treatment` so the call site carries the role.
pub struct TuneParam {
learner_l : LearnerDispatch
learner_m : LearnerDispatch
} derive(Debug)
///|
pub extend TuneParam with @moonbitlang/core/debug.Debug::{to_repr}
///|
/// Build a `TuneParam` from an explicit `(learner_l, learner_m)`
/// pair. Equivalent to the struct literal `{ learner_l, learner_m }`
/// — exposed as a factory so callers don't need to know the field
/// order of an internal struct.
pub fn TuneParam::new(
learner_l : LearnerDispatch,
learner_m : LearnerDispatch,
) -> TuneParam {
{ learner_l, learner_m, }
}
///|
/// Intent-labelled alias of `TuneParam::new` for the **partially
/// linear** families, where the first slot really is the outcome
/// learner `E[y|x]` and the second is the nuisance `E[m|x]`.
///
/// Identical to `TuneParam::new` - it exists so the call site states
/// the intent instead of relying on the field name. Prefer it on
/// `PLR` / `PLIV` / `LPLR` / `PLPR`.
pub fn TuneParam::for_outcome(
learner_l : LearnerDispatch,
learner_m : LearnerDispatch,
) -> TuneParam {
{ learner_l, learner_m, }
}
///|
/// Intent-labelled alias of `TuneParam::new` for the families whose
/// **first slot is the treatment / propensity learner** `E[g|d,x]`
/// and whose second slot is the outcome learner.
///
/// Identical to `TuneParam::new`. Prefer it on `APO` / `APOS` /
/// `IRM` / `IIVM` / `CVAR` / `PQ` / `QTE` / `SSM` and the `DID*`
/// family, so a reader is not left inferring the slot meaning from a
/// field called `learner_l`.
pub fn TuneParam::for_treatment(
learner_g : LearnerDispatch,
learner_m : LearnerDispatch,
) -> TuneParam {
{ learner_l: learner_g, learner_m, }
}
///|
/// Scoring method for `DoubleMLPLR::tune`. All three are based
/// on the OOF outcome-nuisance prediction `l_hat = cross_fit(y | x)`
/// (TUNE_DESIGN.md §4.1: MSE-on-l_hat, the upstream Python default):
///
/// - `MSE`: mean((y - l_hat)^2). Lower is better.
/// - `RMSE`: sqrt(MSE). Lower is better.
/// - `NegMSE`: -MSE. Higher is better (sklearn convention).
///
/// R^2 scoring is deferred to a follow-up release — R^2 needs a
/// defended against degenerate `Var(y) = 0` and isn't needed by
/// the standard PLR tune workflow.
pub enum TuneScoring {
/// Mean squared error; lower is better.
MSE
/// Root mean squared error; lower is better.
RMSE
/// Negative MSE; higher is better (sklearn convention).
NegMSE
} derive(Debug)
///|
pub extend TuneScoring with @moonbitlang/core/debug.Debug::{to_repr}
///|
/// Parse a scoring-method string into a `TuneScoring` enum.
/// Accepted spellings (case-insensitive):
///
/// - `"MSE"` / `"mse"` -> `MSE`
/// - `"RMSE"` / `"rmse"` -> `RMSE`
/// - `"neg-MSE"` / `"neg-mse"` / `"NegMSE"` -> `NegMSE`
///
/// Unknown spellings abort with a descriptive message. Use
/// this when accepting scoring-method names from CLI / JSON
/// configs; pass the enum directly for typed code.
pub fn TuneScoring::parse(s : String) -> TuneScoring {
match s {
"MSE" | "mse" => MSE
"RMSE" | "rmse" => RMSE
"neg-MSE" | "neg-mse" | "NegMSE" | "neg_MSE" => NegMSE
_ => abort("unknown TuneScoring: " + s)
}
}
///|
/// Result of a `DoubleMLPLR::tune` call. Holds the best (by the
/// chosen scoring rule, lower-is-better or higher-is-better per
/// `TuneScoring`) candidate plus the full per-candidate score
/// vector for diagnostics.
///
/// `all_scores[i]` is the raw score for `param_set[i]` (in the
/// grid order, NOT necessarily sorted). Lower-is-better / higher-
/// is-better is the same convention as the chosen `TuneScoring`.
/// For diagnostics where the user wants to compare across scoring
/// rules, they should re-run `tune` with the other rule.
pub struct TuneResult {
/// The candidate with the best score under the chosen rule.
best : TuneParam
/// The best score (raw, in the scoring-rule's own scale).
best_score : Double
/// Per-candidate raw scores, in `param_set` order.
all_scores : Array[Double]
} derive(Debug)
///|
pub extend TuneResult with @moonbitlang/core/debug.Debug::{to_repr}
///|
/// Compute the per-candidate score on a held-out criterion
/// (TUNE_DESIGN.md §4.1: MSE-on-l_hat for outcome-nuisance
/// selection).
///
/// `y` is the outcome vector, `l_hat` is the OOF cross-fit
/// prediction `cross_fit_predict_dispatch(learner_l, x, y, folds)`
/// — i.e. each `l_hat[i]` was produced by a fold that excluded
/// row `i`. Lower MSE means the learner generalizes better
/// out-of-fold.
///
/// Returns a raw, un-sorted score in the chosen metric's natural
/// scale (`MSE` -> mean((y-l_hat)^2), `RMSE` -> sqrt(MSE),
/// `NegMSE` -> -MSE). The caller is responsible for the
/// argmin/argmax selection.
fn tune_score_outcome(
y : Array[Double],
l_hat : Array[Double],
scoring : TuneScoring,
) -> Double {
let n = y.length()
// length pre-check: caller (the `tune` loop) only invokes this
// when `l_hat_c.length() == n`. Adding a `require(...)` here
// would propagate an error type up the call stack and force
// every wrapper to be error-typed too — keep it a plain `let`.
if l_hat.length() != n {
return TUNE_SCORE_FAIL_SENTINEL
}
// raw MSE (un-sorted, lower-is-better)
let mut sse = 0.0
for i = 0; i < n; i = i + 1 {
let r = y[i] - l_hat[i]
sse = sse + r * r
}
let mse = sse / Double::from_int(n)
match scoring {
MSE => mse
RMSE => mse.sqrt()
NegMSE => -mse
}
}
///|
/// Per-candidate raw scores plus the index of the winner.
pub(all) struct TuneGridScore {
/// Raw score per candidate, in the caller's `learners` order.
scores : Array[Double]
/// Index of the winning candidate under the caller's scoring rule.
best_index : Int
}
///|
/// Selection rule: lower score wins, except under `NegMSE` where the
/// sign is already flipped to higher-is-better and the comparison
/// inverts. Private to this file; `tune_score_grid` is the only caller.
fn tune_argmin(scores : Array[Double], scoring : TuneScoring) -> Int {
let mut bi = 0
let mut bv = scores[0]
if scoring is NegMSE {
for i = 1; i < scores.length(); i = i + 1 {
if scores[i] > bv {
bv = scores[i]
bi = i
}
}
} else {
for i = 1; i < scores.length(); i = i + 1 {
if scores[i] < bv {
bv = scores[i]
bi = i
}
}
}
bi
}
///|
/// The estimator-independent half of `tune`: cross-fit every candidate
/// learner against the outcome, score it, and select the winner.
///
/// Extracted in v0.108.0. The five pre-existing `::tune`
/// implementations (`PLR` / `PLIV` / `IRM` / `IIVM` / `APO`) each
/// carried their own byte-identical copy of the scoring loop and of the
/// argmin/argmax selection - roughly 40 duplicated lines apiece,
/// including the `NegMSE` higher-is-better branch that has to be
/// rewritten at every site. That duplication is what let the argmin
/// survive being inverted with the whole suite green (see
/// `expand_v108_test.mbt`).
///
/// `n_folds_tune` folds are drawn fresh from `seed`, deliberately
/// independent of the final fit's own fold schedule, so candidate
/// quality is not selected on the same splits the estimate is later
/// computed from.
///
/// The observation count comes from `y.length()`, NOT from an
/// estimator accessor: six of the estimators this serves (`APOS`
/// `DIDCS` `DIDCrossSection` `DIDMulti` `LPLR` `PLPR`) expose no
/// `n_obs()` at all, and every one of them has `data.y` available
/// instead. A core keyed on `self.n_obs()` would have pushed a new
/// accessor onto six public types purely to serve a helper.
///
/// A candidate whose cross-fit returns a wrong-length vector scores
/// `TUNE_SCORE_FAIL_SENTINEL`, which the selection rule then excludes.
///
/// v0.108.0: the previous version of this comment gave
/// "an RFLearner with 0 trees" as the example, and that was wrong.
/// `RFLearner::fit` raises `PreconditionError` on `n_trees=0` and its
/// own `catch` escalates it to `abort` (`rfl.mbt:404`), so the
/// process dies before any caller sees a vector - let alone a
/// wrong-length one. No learner in the current `LearnerDispatch` set
/// can reach the sentinel: they all return `x.rows()`-length vectors.
/// The guard stays as defence for a future learner that returns a
/// short vector WITHOUT aborting, but it is not a live path today and
/// is not described as one.
///
/// A learner that PANICS cannot be recovered here at all: `abort` is
/// not interceptible, so a panicking candidate aborts the whole
/// process rather than being excluded from the grid.
pub fn tune_score_grid(
x : Matrix,
y : Array[Double],
learners : Array[LearnerDispatch],
n_folds_tune : Int,
seed : Int,
scoring : TuneScoring,
) -> TuneGridScore {
let n = y.length()
let folds_tune = kfold(n, n_folds_tune, seed)
let scores : Array[Double] = Array::make(learners.length(), 0.0)
for i = 0; i < learners.length(); i = i + 1 {
let pred = cross_fit_predict_dispatch(learners[i], x, y, folds_tune)
scores[i] = if pred.length() == n {
tune_score_outcome(y, pred, scoring)
} else {
TUNE_SCORE_FAIL_SENTINEL
}
}
TuneGridScore::{ scores, best_index: tune_argmin(scores, scoring), }
}
///|
/// `DoubleMLPLR::tune` — score a grid of nuisance-learner
/// combinations and re-fit the model with the winner.
///
/// Algorithm (TUNE_DESIGN.md §4):
///
/// 1. Draw a fresh fold schedule `folds_tune` via
/// `kfold(n, n_folds_tune, seed)`. The tune-time folds
/// are deliberately independent of the final-fit
/// `self.n_folds` to avoid information leakage — the
/// final DML estimate uses its own fold schedule for
/// cross-fitting, so sharing tune and fit folds would
/// leak `l_hat` quality into the model-selection step.
/// 2. For each candidate `c` in `param_set`:
/// - `l_hat_c = cross_fit_predict_dispatch(c.learner_l,
/// x, y, folds_tune)`
/// - `score_c = tune_score_outcome(y, l_hat_c, scoring)`
/// - on a learner-level exception, set `score_c = +Inf`
/// (TUNE_DESIGN.md §7: keep tuning, the argmin rule
/// automatically excludes divergent candidates).
/// 3. `c* = argmin_c score_c` (lower-is-better for MSE / RMSE,
/// higher-is-better for NegMSE — `NegMSE` is converted
/// to its lower-is-better form internally).
/// 4. Re-fit the DML model with `learner_l = c*.learner_l,
/// learner_m = c*.learner_m` using the **final-fit** fold
/// schedule (`self.n_folds`). The tuned nuisance
/// predictions from step 1 are NOT reused — each re-fit
/// produces its own OOF cross-fit under `self.n_folds`.
/// 5. Return the re-fitted `DoubleMLPLR`.
///
/// Edge cases (TUNE_DESIGN.md §7):
/// - `param_set.length() == 0` -> abort with a descriptive
/// message.
/// - `param_set.length() == 1` -> skip the scoring loop,
/// re-fit with the single candidate directly (faster).
/// - `n_folds_tune > n_obs` -> abort via the existing
/// `kfold` precondition cascade.
/// - `seed` not set -> defaults to 3141 for reproducibility.
/// - `scoring_method` not in known set -> abort via
/// `TuneScoring::parse`.
///
/// `score` (the DML score: "partialling-out" / "iv-type") is
/// forwarded to the re-fit step unchanged — tune() does not
/// affect the DML score itself, only the nuisance learners.
///
/// Cluster-data (`self.data.is_cluster_data()`) is not
/// supported in v0.58.0 — cluster-aware tuning needs a
/// different fold schedule (`kfold` on unique units) and
/// would inflate this method beyond a single release. Use
/// `DoubleMLPLR::fit(learner_l = best_l, learner_m = best_m)`
/// directly on the cluster-DML path with hand-picked
/// candidates until v0.59.0+ lands cluster-aware tune.
pub fn DoubleMLPLR::tune(
self : DoubleMLPLR,
param_set~ : Array[TuneParam],
scoring_method? : String = "MSE",
n_folds_tune? : Int = 5,
seed? : Int = 3141,
score? : String = "partialling-out",
) -> DoubleMLPLR {
try {
require(param_set.length() > 0)
// `n_folds_tune >= 2` stays at the call site rather than in the
// shared core, and that is deliberate. It is the check that
// carries the most weight - `kfold(n, 1, seed)` yields one fold,
// so the "cross-fit" degenerates to in-sample prediction and every
// candidate gets scored on its own training fit, silently
// rewarding the most overfitting candidate instead of the most
// generalising one. A check that important should be visible at
// every call site, not inferred from a helper. It is also the
// reason the core is not an error-typed function: a polymorphic
// `raise` here would widen the try body's error set and break the
// `catch { PreconditionError::Violated(loc) => ... }` exhaustiveness
// that every `::tune` relies on.
require(n_folds_tune >= 2)
// v0.58.0: cluster-data not yet supported (see file header).
require(!self.data.is_cluster_data())
let scoring = TuneScoring::parse(scoring_method)
// v0.108.0: the scoring loop, the fold draw and the argmin/argmax
// selection moved into `tune_score_grid`. The single-candidate
// fast path went with them: it never actually skipped the scoring
// (it computed one score and hand-built a `TuneResult` by hand),
// so it was a second copy of the result-construction logic that
// could drift from the multi-candidate path rather than a
// shortcut. With one candidate the grid scores it and selects
// index 0, which is the same answer by a shorter route.
let grid = tune_score_grid(
self.data.x,
self.data.y,
param_set.map(fn(p : TuneParam) { p.learner_l }),
n_folds_tune,
seed,
scoring,
)
let best_param = param_set[grid.best_index]
// `best_score` is the raw winning score in the chosen metric's
// natural scale (lower-is-better for MSE / RMSE, higher-is-better
// for NegMSE). Callers compare across scoring rules by re-running
// tune with a different rule.
let result : TuneResult = {
best: best_param,
best_score: grid.scores[grid.best_index],
all_scores: grid.scores,
}
// Final re-fit with the chosen (learner_l, learner_m) under the
// FINAL-FIT fold schedule (`self.n_folds`), NOT under the tune
// folds. The re-fit call produces its own OOF cross-fit and DML
// score.
self.fit(
learner_l=best_param.learner_l,
learner_m=best_param.learner_m,
score~,
tune_result=Some(result),
)
} catch {
PreconditionError::Violated(loc) =>
abort("precondition failed at " + loc.to_string())
}
}
// --- internal helpers (package-private) ------------------------------
///|
/// Sentinel score for a divergent learner in the tune grid
/// (TUNE_DESIGN.md §7: "score_c = +Inf on a learner failure").
/// 1e300 is large enough that any finite MSE / RMSE compares
/// as strictly less, so the argmin rule automatically excludes
/// the divergent candidate. Avoids depending on
/// `@math.inf` / platform-specific NaN/Inf semantics across
/// the 4 backends (wasm, wasm-gc, js each handle IEEE-754
/// infinity differently when serialized through the FFI
/// boundary).
const TUNE_SCORE_FAIL_SENTINEL : Double = 1.0e300