///|
// Memoization layer for the cross-fitting step of DML estimators
// (v0.80.0). When `DoubleMLIRM` is created with `enable_memoize()`,
// repeated `fit()` calls with the same data fingerprint, fold
// split, and nuisance-learner configuration can reuse the cached
// per-observation nuisance predictions instead of re-fitting the
// outcome / propensity learners on every call.
//
// The cache is invalidated when any of the following change:
// - the data (`hash_data(X, y, d, z, cluster_vars)`)
// - the fold split (`fold_split_seed`, `n_folds`, `n_rep`)
// - the observation count (`n_obs`)
// - the nuisance-learner configuration
// (`hash_hyperparams(estimator_kind, ml_g, ml_m)`)
// - the estimator family (`estimator_kind` string)
//
// Scope: only the nuisance-prediction step is memoized. The
// coefficient / SE aggregation from the cached nuisances still
// runs on every `fit()` call (the cached `psi_a` / `psi_b` and
// the post-aggregation `(coef, se)` are recomputed, so changing
// the score / aggregator configuration always takes effect).
//
// Vectorization of the per-fold nuisance fit/predict is deferred
// to v0.81 -- this release only adds the caching layer.
///|
/// Memoization container for a single DML `fit()` call.
///
/// Stored on the estimator struct (e.g. `DoubleMLIRM::fit_cache`).
/// `predictions` is a list of per-observation nuisance prediction
/// arrays whose meaning depends on the estimator kind:
///
/// - `"irm"`: `[g0_hat, g1_hat, m_hat, m_raw]`
/// - `"plr"`: `[g_hat, m_hat]`
/// - `"iivm"`: `[g0_hat, g1_hat, m_hat, r0_hat, r1_hat]`
/// - `"pliv"`: `[g_hat, m_hat]`
/// - `"did"` / `"did_cs"`: `[g_pre, g_post, m_pre, m_post]`
/// - `"lpq"` / `"cvar"` / `"ssm"`: per-estimator nuisance arrays
///
/// `fold_ids[i]` is the test-fold index of observation `i` in
/// the LAST repetition (the one whose nuisances are stored).
///
/// `cluster_ids_hash` (v0.82.0+) is the hash of the
/// `cluster_vars` vector (or 0 for the IID path). The
/// `is_valid(...)` check requires the stored hash to match the
/// caller's hash, so a switch from IID to clustered (or a
/// different cluster partition) invalidates the cache.
pub struct FitCache {
fold_ids : Array[Int]
predictions : Array[Array[Double]]
fold_split_seed : Int
n_folds : Int
n_rep : Int
n_obs : Int
data_hash : UInt64
hyperparams_hash : UInt64
cluster_ids_hash : UInt64
estimator_kind : String
} derive(Debug)
///|
pub extend FitCache with @moonbitlang/core/debug.Debug::{to_repr}
///|
/// Empty cache sentinel. `fold_ids.length() == 0` is the "no
/// cache" marker; `is_valid(...)` returns `false` against this
/// empty cache because the per-call `n_obs` / `n_folds` /
/// hashes can never simultaneously match `0`.
pub fn FitCache::empty() -> FitCache {
{
fold_ids: [],
predictions: [],
fold_split_seed: 0,
n_folds: 0,
n_rep: 0,
n_obs: 0,
data_hash: 0UL,
hyperparams_hash: 0UL,
cluster_ids_hash: 0UL,
estimator_kind: "",
}
}
///|
/// `true` iff the cache holds no observations and was never
/// populated. The opposite of `is_valid(...)` against the
/// expected inputs.
pub fn FitCache::is_empty(self : FitCache) -> Bool {
self.fold_ids.length() == 0
}
///|
/// Check whether the stored cache matches the requested
/// configuration. Returns `true` only if every dimension of the
/// fingerprint matches the current call -- otherwise the caller
/// must fall back to the standard (non-cached) cross-fit.
///
/// `cluster_ids_hash` (v0.82.0+) is the hash of the
/// `cluster_vars` vector (or 0 for the IID path). Pass the same
/// hash you used when the cache was populated; the caller's
/// hash is compared to the stored one to detect a switch from
/// IID to clustered (or a different cluster partition).
pub fn FitCache::is_valid(
self : FitCache,
fold_split_seed : Int,
n_folds : Int,
n_rep : Int,
n_obs : Int,
data_hash : UInt64,
hyperparams_hash : UInt64,
cluster_ids_hash : UInt64,
estimator_kind : String,
) -> Bool {
self.fold_ids.length() == n_obs &&
self.fold_split_seed == fold_split_seed &&
self.n_folds == n_folds &&
self.n_rep == n_rep &&
self.n_obs == n_obs &&
self.data_hash == data_hash &&
self.hyperparams_hash == hyperparams_hash &&
self.cluster_ids_hash == cluster_ids_hash &&
self.estimator_kind == estimator_kind
}
///|
/// FoldMix64 step: `(h ^ v) * FNV_PRIME_64`. The FNV-1a style
/// multiplication keeps the avalanche fast and the resulting
/// hash well-distributed. `FNV_PRIME_64` is the canonical 64-bit
/// FNV prime (1099511628211).
fn fold_mix(h : UInt64, v : UInt64) -> UInt64 {
let prime : UInt64 = 0x00000100000001B3UL
(h ^ v) * prime
}
///|
/// Mix a Double into the hash. MoonBit's `Double::to_uint64`
/// and `to_int64` truncate the value (not a bit-cast), so
/// using them directly would collide e.g. `0.999` with `0.0`.
/// We use `Double::to_string()` instead: the MoonBit runtime
/// emits a canonical scientific-notation string for each IEEE
/// bit pattern, so two different Doubles produce two different
/// strings and therefore two different hashes. The string is
/// slow to produce (relative to a bit cast), but the cost is
/// amortized across a single `fit()` call (the hash is
/// computed once per fit, not per access), so this stays
/// acceptable for memoization.
fn fold_mix_double(h : UInt64, v : Double) -> UInt64 {
let s = v.to_string()
let mut acc = h
for c in s {
acc = fold_mix_int(acc, c.to_int())
}
acc
}
///|
/// Mix an Int into the hash. Two's complement bit pattern gives
/// a stable mapping for all integer values (including negatives).
fn fold_mix_int(h : UInt64, v : Int) -> UInt64 {
fold_mix(h, v.to_uint64())
}
///|
/// 64-bit content hash of the data passed to a DML estimator.
/// Combines the data dimensions, a sampled slice of `X` (first
/// 10 + last 10 rows, first column), and the full content of
/// `y`, `d`, `z`, and `cluster_vars`.
///
/// The hash is NOT cryptographic -- it is designed to detect
/// accidental input mutations (row reordering, scaling, value
/// edits, dimension changes) at O(n) cost. A malicious caller
/// could craft two different datasets that share the same hash;
/// that is acceptable for a memoization invalidation key.
pub fn hash_data(
x : Matrix,
y : Array[Double],
d : Array[Double],
z? : Array[Double] = [],
cluster_vars? : Array[Int] = [],
) -> UInt64 {
let mut h : UInt64 = 0xcbf29ce484222325UL // FNV-1a 64-bit offset basis
// Dimensions + vector lengths first so a dimension change
// invalidates the cache even before sampling any values.
h = fold_mix_int(h, x.rows())
h = fold_mix_int(h, x.cols())
h = fold_mix_int(h, y.length())
h = fold_mix_int(h, d.length())
h = fold_mix_int(h, z.length())
h = fold_mix_int(h, cluster_vars.length())
// X sample: first 10 + last 10 rows, column 0 only. Cheap
// (2*10 = 20 cell accesses) and catches the common cases of
// row scaling / column edits / append-prepend mutations.
let n_rows = x.rows()
let n_cols = x.cols()
if n_rows > 0 && n_cols > 0 {
let first_n = if n_rows < 10 { n_rows } else { 10 }
for i = 0; i < first_n; i = i + 1 {
h = fold_mix_double(h, x.get(i, 0))
}
if n_rows > 10 {
let last_start = n_rows - 10
for i = 0; i < 10; i = i + 1 {
h = fold_mix_double(h, x.get(last_start + i, 0))
}
}
}
// Full content of the per-observation vectors.
for v in y {
h = fold_mix_double(h, v)
}
for v in d {
h = fold_mix_double(h, v)
}
for v in z {
h = fold_mix_double(h, v)
}
for v in cluster_vars {
h = fold_mix_int(h, v)
}
h
}
///|
/// Hyperparameter fingerprint. Combines the estimator family
/// name, the two nuisance-learner dispatch tags (`ml_g` / `ml_m`),
/// and the propensity-clip tolerance (only used by IRM but
/// included unconditionally for cache-key uniformity).
///
/// Each `LearnerDispatch` variant is matched on a tag string
/// (`"linear_regression"` / `"logistic_regression"` / `"constant"`
/// / `"noop"` / `"random_forest"` / `"gradient_boosting"`). The
/// variant's tunables (e.g. RF `n_trees` / `max_depth`) are NOT
/// folded into the key in this release -- the same learner tag
/// re-used with different tunables will hit the cache and yield
/// stale predictions. Users who change tunables across calls
/// must call `DoubleMLIRM::clear_cache()` first. This trade-off
/// keeps the key short and matches the conventional "same
/// learner name = same learner" contract used by upstream
/// `doubleml-for-py`.
pub fn hash_hyperparams(
estimator_kind : String,
ml_g : LearnerDispatch,
ml_m : LearnerDispatch,
propensity_clip : Double,
) -> UInt64 {
let mut h : UInt64 = 0x84222325cbf29ce4UL
h = fold_mix_int(h, estimator_kind.length())
// Fold the bytes of estimator_kind as Int hashes -- this gives
// a length-aware, content-aware mix without depending on
// `String` -> `UInt64` direct conversion.
for c in estimator_kind {
h = fold_mix_int(h, c.to_int())
}
h = fold_mix(h, learner_dispatch_tag(ml_g))
h = fold_mix(h, learner_dispatch_tag(ml_m))
h = fold_mix_double(h, propensity_clip)
h
}
///|
/// Map a `LearnerDispatch` to a stable 64-bit tag. The tag is
/// derived from the variant's name string so two learners of the
/// same kind produce the same tag. Tunables are intentionally
/// NOT folded in -- see `hash_hyperparams` for the rationale.
fn learner_dispatch_tag(l : LearnerDispatch) -> UInt64 {
let kind : String = match l {
LinearRegression(_) => "linear_regression"
LogisticRegression(_) => "logistic_regression"
Constant(_) => "constant"
Noop(_) => "noop"
RandomForest(_) => "random_forest"
GradientBoosting(_) => "gradient_boosting"
}
let mut h : UInt64 = 0x9e3779b97f40b847UL
for c in kind {
h = fold_mix_int(h, c.to_int())
}
h
}
///|
/// Build a `FitCache` from the LAST repetition's fold assignment
/// and nuisance predictions. `predictions` is interpreted by the
/// estimator (see the `FitCache` docstring for the per-kind
/// layout).
pub fn FitCache::from_fit(
fold_ids : Array[Int],
predictions : Array[Array[Double]],
fold_split_seed : Int,
n_folds : Int,
n_rep : Int,
n_obs : Int,
data_hash : UInt64,
hyperparams_hash : UInt64,
cluster_ids_hash : UInt64,
estimator_kind : String,
) -> FitCache {
{
fold_ids,
predictions,
fold_split_seed,
n_folds,
n_rep,
n_obs,
data_hash,
hyperparams_hash,
cluster_ids_hash,
estimator_kind,
}
}
///|
/// Compute the cache-key hash of a cluster-ids vector.
/// v0.82.0+: an empty vector hashes to 0 (the IID sentinel);
/// a non-empty vector hashes via FNV-style fold over its
/// entries. Two structurally-different cluster partitions
/// produce different hashes so the cache invalidates on a
/// cluster redefinition (e.g. a switch from individual-level
/// clustering to individual-by-time or a different
/// treatment-period partition).
pub fn hash_cluster_ids(cluster_ids : Array[Int]) -> UInt64 {
if cluster_ids.length() == 0 {
return 0UL
}
let mut h : UInt64 = 0xcbf29ce484222325UL
// Length first so a length change invalidates the cache even
// before sampling any value.
h = fold_mix_int(h, cluster_ids.length())
for v in cluster_ids {
h = fold_mix_int(h, v)
}
h
}