///|
/// `TuneParam` — one row of the candidate grid passed to
/// `DoubleMLPLR::tune`. Holds a concrete nuisance-learner
/// combination (one `learner_l`, one `learner_m`) 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`.
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 }
}
///|
/// 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
}
}
///|
/// `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)
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)
// Single-entry grid: skip the scoring loop, re-fit directly
// (TUNE_DESIGN.md §7: "param_set.length() == 1: skip tune,
// return re-fit"). Re-bounded control flow keeps the loop
// body uniform across all paths.
if param_set.length() == 1 {
let c0 = param_set[0]
// No score comparison possible with one candidate, but
// still record a TuneResult so callers can see what was
// chosen. `best_score` is the in-sample MSE on `l_hat`
// computed under `folds_tune` (purely informational —
// tune() did not run an argmin over multiple candidates).
let n = self.n_obs()
let folds_tune = kfold(n, n_folds_tune, seed)
let l_hat_0 = cross_fit_predict_dispatch(
c0.learner_l, self.data.x, self.data.y, folds_tune,
)
let score_0 = if l_hat_0.length() == n {
tune_score_outcome(self.data.y, l_hat_0, scoring)
} else {
TUNE_SCORE_FAIL_SENTINEL
}
let result : TuneResult = {
best: c0,
best_score: score_0,
all_scores: [score_0],
}
return self.fit(
learner_l=c0.learner_l,
learner_m=c0.learner_m,
score~,
tune_result=Some(result),
)
}
let n = self.n_obs()
let folds_tune = kfold(n, n_folds_tune, seed)
let scores : Array[Double] = Array::make(param_set.length(), 0.0)
for i = 0; i < param_set.length(); i = i + 1 {
let c = param_set[i]
// Per-candidate cross-fit under `folds_tune`. A defensive
// length check on `l_hat_c` covers the divergence case
// (TUNE_DESIGN.md §7): if a learner returns a wrong-size
// array (e.g. an RFLearner with 0 trees), the score is
// set to +Inf so the argmin rule excludes the bad
// candidate. True exception-based divergent-learner
// recovery is deferred to v0.59.0+ when `Learner::predict`
// becomes Result-typed.
let l_hat_c = cross_fit_predict_dispatch(
c.learner_l, self.data.x, self.data.y, folds_tune,
)
scores[i] = if l_hat_c.length() == n {
tune_score_outcome(self.data.y, l_hat_c, scoring)
} else {
TUNE_SCORE_FAIL_SENTINEL
}
}
// argmin over `scores`. For NegMSE (higher-is-better) we
// negate so the same argmin rule applies uniformly.
let best_idx = if scoring is NegMSE {
let mut bi = 0
let mut bv = scores[0]
for i = 1; i < scores.length(); i = i + 1 {
if scores[i] > bv {
bv = scores[i]
bi = i
}
}
bi
} else {
let mut bi = 0
let mut bv = scores[0]
for i = 1; i < scores.length(); i = i + 1 {
if scores[i] < bv {
bv = scores[i]
bi = i
}
}
bi
}
let best_param = param_set[best_idx]
// Build the TuneResult to record on the re-fitted model.
// `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: scores[best_idx],
all_scores: scores,
}
// Final re-fit with the chosen (learner_l, learner_m) under
// the FINAL-FIT fold schedule (`self.n_folds`), NOT under
// `folds_tune`. 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