///|
/// Best linear predictor of an orthogonal signal on a supplied basis.
///
/// `cov_type` selects the standard-error convention. Valid values:
///   - `"HC0"` (default): White's heteroskedasticity-consistent
///     sandwich SE — robust to arbitrary residual heteroskedasticity.
///     Matches the upstream `doubleml.utils.blp` call to
///     `statsmodels.OLS(cov_type='HC0')`.
///   - `"nonrobust"`: classic homoskedastic OLS SE
///     (`sigma^2 * (X^T X)^{-1}` with `sigma^2 = RSS / (n - p)`),
///     which is only valid under the homogeneous-error assumption.
///
/// REVIEW L5: the field is currently stored on the struct for forward
/// compatibility (post-fit introspection), but is not used outside
/// `fit`. Construct via `DoubleMLBLP::new` to validate the value.
///
/// v0.63.0+: `ml_g?` field accepts a `LearnerDispatch`. The
/// `coef` / `(X^T X)^{-1}` machinery is always the closed-form OLS
/// path (BLP is by definition the OLS projection); the dispatch
/// learner only affects the residual / RSS computation. Defaults
/// to `LearnerDispatch::linear_regression()` which preserves the
/// v0.62.2 byte-equality behaviour for the standard fit.
pub struct DoubleMLBLP {
  basis : Matrix
  orth_signal : Array[Double]
  cov_type : String
  ml_g : LearnerDispatch
  coef : Array[Double]
  se : Array[Double]
  fitted : Bool
  // v0.19.0+: post-fit summary statistics used by
  // `GainStatsSource::from_blp` to auto-populate the
  // sensitivity parameter benchmarks. Initialised to
  // zero in `new`; filled in by `fit`.
  n_obs : Int
  rss : Double
  var_y : Double
  // v0.64.0+: per-observation residual vector at the fitted
  // OLS projection. Length `n_obs`. Persisted for the
  // multiplier bootstrap.
  residuals : Array[Double]
  // v0.64.0+: multiplier bootstrap state. `boot_t_stat` is
  // a flat `[n_rep_boot * (p + 1)]` array of t-statistics
  // (row-major by rep, then by coef — same layout as the
  // rest of the package). Populated by `bootstrap(...)`;
  // empty until then.
  boot_t_stat : Array[Double]
  boot_method : String
  n_rep_boot : Int
  boot_seed : Int
  // v0.83.0+: memoization state. `memoize_enabled` is the
  // user-facing switch (false by default to preserve v0.82.0
  // behavior bit-for-bit). When true, `fit()` caches the
  // single-pass `(coef, se, n_obs, rss, var_y, residuals)`
  // plus a placeholder fold_ids vector in `fit_cache` and
  // reuses them on the next call when the basis, orthogonal
  // signal, learner fingerprint, and covariance type are
  // unchanged. Mirrors the IRM / PLR / CVAR / SSM plumbing.
  memoize_enabled : Bool
  fit_cache : FitCache
} derive(Debug)

///|
pub extend DoubleMLBLP with @moonbitlang/core/debug.Debug::{to_repr}

///|
pub fn DoubleMLBLP::new(
  basis : Matrix,
  orth_signal : Array[Double],
  cov_type? : String = "HC0",
  ml_g? : LearnerDispatch = LearnerDispatch::linear_regression(),
) -> DoubleMLBLP {
  try {
    require(basis.rows() == orth_signal.length())
    require(cov_type == "HC0" || cov_type == "nonrobust")
    {
      basis,
      orth_signal,
      cov_type,
      ml_g,
      coef: [],
      se: [],
      fitted: false,
      n_obs: 0,
      rss: 0.0,
      var_y: 0.0,
      residuals: Array::make(basis.rows(), 0.0),
      boot_t_stat: [],
      boot_method: "",
      n_rep_boot: 0,
      boot_seed: 0,
      // v0.83.0+: default memoize off so v0.82.0 callers see
      // byte-identical fit() output. Enable explicitly via
      // `.enable_memoize()` for caching.
      memoize_enabled: false,
      fit_cache: FitCache::empty(),
    }
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
pub fn DoubleMLBLP::fit(
  self : DoubleMLBLP,
  ml_g? : LearnerDispatch = self.ml_g,
) -> DoubleMLBLP {
  // v0.83.0+: memoize check. BLP is a single-pass OLS
  // projection (no `n_folds` / `n_rep`), so the cache stores
  // the post-fit `(coef, se, residuals, n_obs, rss, var_y)`
  // plus a length-`n_obs` placeholder `fold_ids` (BLP does
  // not partition rows into folds; the cache layout still
  // expects a `Array[Int]` of length `n_obs` so the
  // `FitCache::is_valid` check has a consistent shape).
  let n_obs = self.orth_signal.length()
  let memoize = self.memoize_enabled
  let data_hash : UInt64 = if memoize {
    // BLP has no `d` / `z` / `cluster_vars`; pass empty
    // arrays so the hash reflects only the basis +
    // orth_signal content.
    hash_data(self.basis, self.orth_signal, [])
  } else {
    0UL
  }
  let hparams_hash : UInt64 = if memoize {
    // BLP's only nuisance learner is `ml_g`. We pass `ml_g`
    // twice (once for the `ml_g` slot, once as a placeholder
    // for the `ml_m` slot BLP doesn't carry) and use a
    // `propensity_clip` proxy derived from `cov_type` --
    // 0.0 for "HC0", 1.0 for "nonrobust" -- so the
    // hyperparams hash also invalidates on a covariance-type
    // switch.
    let cov_clip = if self.cov_type == "HC0" { 0.0 } else { 1.0 }
    hash_hyperparams("blp", ml_g, ml_g, cov_clip)
  } else {
    0UL
  }
  let cluster_hash : UInt64 = if memoize {
    // BLP data has no `cluster_vars`; pass an empty vector
    // so the hash stays 0 across calls (matches the IID
    // sentinel).
    hash_cluster_ids([])
  } else {
    0UL
  }
  let cache_hit = memoize &&
    self.fit_cache.is_valid(
      self.cov_type.length(),
      // BLP has no `seed` / `n_folds`; use `cov_type.length()`
      // (3 for "HC0", 9 for "nonrobust") as a fingerprint so a
      // covariance-type change invalidates the cache. The
      // `n_folds` slot is fed the basis row count as a
      // placeholder so the cache key is still dimensional.
      self.basis.rows(),
      1,
      n_obs,
      data_hash,
      hparams_hash,
      cluster_hash,
      "blp",
    )
  if cache_hit {
    // Reuse the cached OLS projection. The returned struct
    // mirrors the constructor's "fitted" state but with
    // `fitted = true` and the cached arrays.
    let preds = self.fit_cache.predictions
    let c : Array[Double] = preds[0]
    let s_arr : Array[Double] = preds[1]
    let cached_rss : Double = preds[2][0]
    let cached_var_y : Double = preds[3][0]
    let cached_residuals : Array[Double] = preds[4]
    return {
      basis: self.basis,
      orth_signal: self.orth_signal,
      cov_type: self.cov_type,
      ml_g,
      coef: c,
      se: s_arr,
      fitted: true,
      n_obs,
      rss: cached_rss,
      var_y: cached_var_y,
      residuals: cached_residuals,
      boot_t_stat: [],
      boot_method: "",
      n_rep_boot: 0,
      boot_seed: 0,
      memoize_enabled: self.memoize_enabled,
      fit_cache: self.fit_cache,
    }
  }
  // v0.63.0+: route through `LearnerDispatch` so the per-fit
  // `ml_g` override reaches the BLP fit. Defaults preserve
  // v0.62.2 byte-equality (default `LearnerDispatch::linear_regression()`
  // matches the hardcoded `LinearRegression::new()` path).
  // The closed-form `(X^T X)^{-1}` machinery stays on the
  // OLS path (BLP is by construction the OLS projection);
  // the dispatch learner only affects the residual / RSS
  // computation.
  let ols_model = LinearRegression::new().fit(self.basis, self.orth_signal)
  let c = ols_model.coefficients()
  let p = c.length()
  let s = Array::make(p, 0.0)
  // Per-coefficient SE. Two flavours:
  //   - HC0 (sandwich, default): `cov_jj = sum_i ((M[j,:] x_i)^2 * e_i^2)`
  //     matches upstream `statsmodels.OLS(cov_type='HC0')`. Robust
  //     to arbitrary heteroskedasticity.
  //   - nonrobust: `cov_jj = sigma^2 * (X^T X)^{-1}_{jj}` with
  //     `sigma^2 = RSS / (n - p)` (homogeneous-error assumption).
  // Per-coefficient SE (Bug #5 fix); HC0 default matches
  // upstream `statsmodels.OLS(cov_type='HC0')`.
  // v0.63.0+: route through `fit_predict_one_dispatch` so the
  // per-fit `ml_g` override drives the residual computation.
  // For the default `LearnerDispatch::linear_regression()` this
  // matches the v0.62.2 `model.predict(self.basis)` path
  // byte-for-byte (the dispatch wrapper calls `lr.fit(basis,
  // orth_signal).predict(basis)` internally).
  let pred = fit_predict_one_dispatch(
    ml_g,
    self.basis,
    self.orth_signal,
    self.basis,
  )
  // v0.83.0+: vectorise the residual / RSS computation. Use
  // `vector_subtract` for the per-observation residual and
  // `vector_multiply` for the squared-residual, then
  // `mean * n` for the sum (no public `sum` helper; this
  // matches the upstream convention of a two-pass mean).
  let residuals : Array[Double] = vector_subtract(self.orth_signal, pred)
  let sq_resid : Array[Double] = vector_multiply(residuals, residuals)
  let rss = mean(sq_resid) * n_obs.to_double()
  // `var_y` = variance of the orthogonal signal (the BLP's
  // "outcome" variable). Computed as a one-pass Welford
  // would be slightly more accurate, but a two-pass mean
  // is fine for our purposes.
  let mut mean_y = 0.0
  for i = 0; i < n_obs; i = i + 1 {
    mean_y = mean_y + self.orth_signal[i]
  }
  mean_y = mean_y / n_obs.to_double()
  let mut ss_y = 0.0
  for i = 0; i < n_obs; i = i + 1 {
    let d = self.orth_signal[i] - mean_y
    ss_y = ss_y + d * d
  }
  let var_y = ss_y / n_obs.to_double()
  let cov_diag = if self.cov_type == "HC0" {
    ols_model.sandwich_se(self.basis, self.orth_signal)
  } else {
    let sigma2 = rss / (n_obs.to_double() - p.to_double())
    ols_model.covariance_diagonal(sigma2)
  }
  for j = 0; j < p; j = j + 1 {
    s[j] = cov_diag[j].sqrt()
  }
  // v0.83.0+: when memoize is on and the cache missed, write
  // the freshly-computed OLS projection to the cache. The
  // cache layout stores `coef` and `se` as length-`p`
  // arrays, `rss` and `var_y` as length-1 arrays (so they
  // fit the `Array[Array[Double]]` shape of `FitCache`),
  // `residuals` as a length-`n_obs` array, and a length-
  // `n_obs` placeholder `fold_ids` (BLP does not partition
  // rows into folds).
  let next_cache = if memoize {
    FitCache::from_fit(
      Array::make(n_obs, 0),
      [c, s, Array::make(1, rss), Array::make(1, var_y), residuals],
      self.cov_type.length(),
      self.basis.rows(),
      1,
      n_obs,
      data_hash,
      hparams_hash,
      cluster_hash,
      "blp",
    )
  } else {
    FitCache::empty()
  }
  {
    basis: self.basis,
    orth_signal: self.orth_signal,
    cov_type: self.cov_type,
    ml_g,
    coef: c,
    se: s,
    fitted: true,
    n_obs,
    rss,
    var_y,
    residuals,
    boot_t_stat: [],
    boot_method: "",
    n_rep_boot: 0,
    boot_seed: 0,
    memoize_enabled: self.memoize_enabled,
    fit_cache: next_cache,
  }
}

///|
/// v0.83.0+: turn on memoization for subsequent `fit()` calls.
/// When enabled, `fit()` will cache the OLS projection
/// `(coef, se, residuals, n_obs, rss, var_y)` and skip the
/// `LinearRegression::fit` + `sandwich_se` path on a repeat
/// call whose basis + orthogonal signal + learner + covariance
/// type are all unchanged. Mirrors the IRM / PLR / CVAR / SSM
/// plumbing.
///
/// Default is OFF. When OFF, every `fit()` call runs the full
/// OLS projection and the cache is neither read nor written,
/// so v0.82.0 callers see byte-identical output.
pub fn DoubleMLBLP::enable_memoize(self : DoubleMLBLP) -> DoubleMLBLP {
  { ..self, memoize_enabled: true, }
}

///|
/// v0.83.0+: turn off memoization. Same immutability contract
/// as `enable_memoize()`. After this, `fit()` will not read or
/// write the cache; the existing `fit_cache` is preserved on
/// the returned struct (call `clear_cache()` to drop it).
pub fn DoubleMLBLP::disable_memoize(self : DoubleMLBLP) -> DoubleMLBLP {
  { ..self, memoize_enabled: false, }
}

///|
/// v0.83.0+: drop any cached OLS projection. Forces the next
/// `fit()` to recompute from scratch.
pub fn DoubleMLBLP::clear_cache(self : DoubleMLBLP) -> DoubleMLBLP {
  { ..self, fit_cache: FitCache::empty(), }
}

///|
/// v0.83.0+: `true` iff `fit_cache` holds at least one cached
/// observation (i.e. at least one prior `fit()` call with
/// `memoize_enabled = true` has populated the cache). Note
/// that the cache may still be stale relative to the current
/// basis / learner configuration -- check `memoize_enabled`
/// before assuming a cache hit.
pub fn DoubleMLBLP::has_cache(self : DoubleMLBLP) -> Bool {
  !self.fit_cache.is_empty()
}

///|
pub fn DoubleMLBLP::coef(self : DoubleMLBLP) -> Array[Double] {
  try {
    require(self.fitted)
    self.coef
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
pub fn DoubleMLBLP::se(self : DoubleMLBLP) -> Array[Double] {
  try {
    require(self.fitted)
    self.se
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.19.0+: sample size used by the BLP fit.
pub fn DoubleMLBLP::n_obs(self : DoubleMLBLP) -> Int {
  try {
    require(self.fitted)
    self.n_obs
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.64.0+: multiplier bootstrap for `DoubleMLBLP`. BLP is
/// the OLS projection of `orth_signal` onto the augmented
/// design `[1, basis]`; per-observation IF is
/// `psi[i, j] = M[j, :] @ xa_i * e_i` where
/// `M = (Xa^T Xa + ridge I)^{-1}` and `xa_i = [1, basis[i, :]]`.
/// `se_flat` is the per-coefficient SE (HC0 sandwich diagonal
/// under `cov_type = "HC0"`, `sqrt(sigma^2 * (X^T X)^{-1}_{jj})`
/// under `cov_type = "nonrobust"`); the helper skips a
/// coefficient when `se[j] = 0`. Routes through the shared
/// `generic_bootstrap_ols_per_coef` helper (see
/// `bootstrap_helper.mbt`).
///
/// `method_name` selects the multiplier distribution:
/// `"normal"` (default), `"Bayes"`, `"wild"`. `seed` defaults
/// to `2024`; `n_rep_boot` defaults to `500`.
///
/// `boot_t_stat` is a flat `[n_rep_boot * (p + 1)]` array of
/// t-statistics (row-major by rep, then by coef). Calling on
/// an un-fit model aborts via `PreconditionError`.
pub fn DoubleMLBLP::bootstrap(
  self : DoubleMLBLP,
  method_name? : String = "normal",
  n_rep_boot? : Int = 500,
  seed? : Int = 2024,
) -> DoubleMLBLP {
  try {
    require(self.fitted)
    require(
      method_name == "normal" || method_name == "Bayes" || method_name == "wild",
    )
    require(n_rep_boot >= 2)
    // Recompute (Xa^T Xa + ridge I)^{-1} from the augmented
    // design so the IF rows `psi[i, j] = M[j, :] @ xa_i * e_i`
    // are consistent with `LinearRegression::fit` (which
    // augments with an intercept internally and applies the
    // same ridge regularisation to the normal equations).
    let xa = augment_with_intercept(self.basis)
    let xa_t = xa.transpose()
    let xtx = matmul(xa_t, xa)
    let xtx_aug = add_ridge(xtx, 1.0e-10)
    let xtx_inv = inv_spd(xtx_aug)
    let boot_t_stat = generic_bootstrap_ols_per_coef(
      xa,
      self.residuals,
      xtx_inv,
      self.se,
      method_name,
      n_rep_boot,
      seed,
    ) catch {
      BootstrapMethodError::UnknownMethod(m) =>
        abort(
          "draw_bootstrap_weights: unknown method (set in DoubleMLBLP::bootstrap): " +
          m,
        )
    }
    {
      ..self,
      boot_t_stat,
      boot_method: method_name,
      n_rep_boot,
      boot_seed: seed,
    }
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.19.0+: residual sum of squares from the BLP fit.
/// Equals `sum_i (orth_signal[i] - basis[i] @ coef)^2`.
pub fn DoubleMLBLP::rss(self : DoubleMLBLP) -> Double {
  try {
    require(self.fitted)
    self.rss
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.69.0+: per-coefficient sensitivity analysis for `DoubleMLBLP`.
///
/// BLP is the closed-form OLS projection of `orth_signal` onto the
/// augmented design `Xa = [1, basis]` (intercept + basis columns).
/// Per-observation influence-function rows for coefficient `j` are
///
///   psi_a[j][i] = M[j, :] @ xa_i
///
/// where `M = (Xa^T Xa + ridge I)^{-1}` is the OLS precision matrix
/// and `xa_i` is the augmented row `i`. The full IF for coef `j`
/// decomposes as `psi[i, j] = psi_a[j][i] * residuals[i]` (matches
/// the v0.64.0 `bootstrap` IF rows used by
/// `generic_bootstrap_ols_per_coef`).
///
/// Sensitivity for each coef follows the shared
/// `irm_style_sensitivity` decomposition with:
///   - `theta = coef[j]`
///   - `residuals = self.residuals` (the OLS residuals from `fit`)
///   - `psi_a = psi_a[j]`
///   - `cf_y / cf_d` default to `0.05` (matching the rest of the
///     v0.66.0-onward sensitivity family).
///
/// Returns an `Array[SensitivityResult]` of length `coef.length()`
/// (= `n_features + 1`, including the intercept at index 0).
/// Calling on an un-fit model aborts via `PreconditionError`.
pub fn DoubleMLBLP::sensitivity_analysis(
  self : DoubleMLBLP,
  cf_y? : Double = 0.05,
  cf_d? : Double = 0.05,
) -> Array[SensitivityResult] raise {
  require(self.fitted)
  let n_obs = self.orth_signal.length()
  let n_coef = self.coef.length()
  // Reconstruct the augmented design + precision matrix (matches
  // the v0.64.0 `bootstrap` IF-row computation so the IRM-style
  // decomposition here stays consistent with the bootstrap path).
  let xa = augment_with_intercept(self.basis)
  let p_aug = xa.ncols
  let xa_t = xa.transpose()
  let xtx = matmul(xa_t, xa)
  let xtx_aug = add_ridge(xtx, 1.0e-10)
  let xtx_inv = inv_spd(xtx_aug)
  let results : Array[SensitivityResult] = Array::make(n_coef, {
    rv: 0.0,
    sigma2: 0.0,
    nu2: 0.0,
    cf_y: 0.0,
    cf_d: 0.0,
    max_bias: 0.0,
  })
  for j = 0; j < n_coef; j = j + 1 {
    let psi_a_j : Array[Double] = Array::make(n_obs, 0.0)
    for i = 0; i < n_obs; i = i + 1 {
      // psi_a[j][i] = M[j, :] @ xa_i = sum_k xtx_inv[j, k] * xa[i, k]
      let mut dot = 0.0
      for k = 0; k < p_aug; k = k + 1 {
        dot = dot + xtx_inv.data[j * p_aug + k] * xa.data[i * p_aug + k]
      }
      psi_a_j[i] = dot
    }
    results[j] = irm_style_sensitivity(
      self.coef[j],
      self.residuals,
      psi_a_j,
      cf_y,
      cf_d,
    )
  }
  results
}

///|
/// v0.74.0+: cluster-robust analogue of
/// `DoubleMLBLP::sensitivity_analysis`. Same residual
/// (`self.residuals`) and same per-coef `psi_a_j` (the
/// `M[j, :] @ xa_i` Riesz-representer rows from the OLS
/// precision matrix) as the IID path; only the variance /
/// bias computation is cluster-aware. The augmented
/// design `xa` and the precision `xtx_inv` are recomputed
/// identically so the per-coef `psi_a_j` rows match the
/// bootstrap IF rows exactly.
///
/// `DoubleMLBLP` has no `data` field (just `basis` +
/// `orth_signal`); the user must pass `cluster_ids`
/// explicitly. `cluster_ids.length()` must equal
/// `self.basis.rows()` (= `self.orth_signal.length()`).
///
/// Returns an `Array[SensitivityResult]` of length
/// `self.coef.length()` (= `n_features + 1`, including the
/// intercept at index 0). Calling on an un-fit model aborts
/// via `PreconditionError`.
pub fn DoubleMLBLP::sensitivity_analysis_cluster(
  self : DoubleMLBLP,
  cluster_ids : Array[Int],
  cf_y? : Double = 0.05,
  cf_d? : Double = 0.05,
) -> Array[SensitivityResult] raise {
  require(self.fitted)
  let n_obs = self.orth_signal.length()
  let n_coef = self.coef.length()
  require(cluster_ids.length() == n_obs)
  // Reconstruct the augmented design + precision matrix
  // (matches the v0.64.0 `bootstrap` IF-row computation so
  // the IRM-style decomposition here stays consistent with
  // the bootstrap path).
  let xa = augment_with_intercept(self.basis)
  let p_aug = xa.ncols
  let xa_t = xa.transpose()
  let xtx = matmul(xa_t, xa)
  let xtx_aug = add_ridge(xtx, 1.0e-10)
  let xtx_inv = inv_spd(xtx_aug)
  // Build per-coef `psi_a` rows once, then loop the cluster
  // helper per coef.
  let psi_a_all : Array[Array[Double]] = Array::make(n_coef, [])
  for j = 0; j < n_coef; j = j + 1 {
    let psi_a_j : Array[Double] = Array::make(n_obs, 0.0)
    for i = 0; i < n_obs; i = i + 1 {
      // psi_a[j][i] = M[j, :] @ xa_i = sum_k xtx_inv[j, k] * xa[i, k]
      let mut dot = 0.0
      for k = 0; k < p_aug; k = k + 1 {
        dot = dot + xtx_inv.data[j * p_aug + k] * xa.data[i * p_aug + k]
      }
      psi_a_j[i] = dot
    }
    psi_a_all[j] = psi_a_j
  }
  // For BLP the residuals are shared across coefs (they are
  // the OLS projection residuals of `orth_signal` on
  // `xa`). Replicate into n_coef rows for the multi helper.
  let residuals_arr : Array[Array[Double]] = Array::make(n_coef, [])
  for j = 0; j < n_coef; j = j + 1 {
    residuals_arr[j] = self.residuals.copy()
  }
  irm_style_sensitivity_cluster_multi(
    self.coef,
    residuals_arr,
    psi_a_all,
    cluster_ids,
    cf_y,
    cf_d,
  )
}

///|
/// v0.19.0+: variance of the orthogonal signal (the BLP's
/// "outcome" variable). Computed as a population variance
/// (divisor `n`, not `n - 1`).
pub fn DoubleMLBLP::var_y(self : DoubleMLBLP) -> Double {
  try {
    require(self.fitted)
    self.var_y
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.22.0+: the orthogonal signal array (the BLP's
/// "outcome" variable). Length `n_obs`. Used by
/// `GainStatsSource::from_blp_cv` to compute the
/// cross-fit residual variance.
pub fn DoubleMLBLP::orth_signal(self : DoubleMLBLP) -> Array[Double] {
  try {
    require(self.fitted)
    self.orth_signal
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.22.0+: the basis matrix (the BLP's "design
/// matrix"). Shape `n_obs x p_features`. Used by
/// `GainStatsSource::from_blp_cv` to refit the
/// BLP on each fold's training subset.
pub fn DoubleMLBLP::basis(self : DoubleMLBLP) -> Matrix {
  try {
    require(self.fitted)
    self.basis
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// Fitted values `basis_aug @ coef` for each observation, length
/// `n_obs`. Matches the upstream `DoubleMLBLP.predictions` /
/// `predict` semantic — the orthogonal-signal prediction under
/// the BLP coefficient vector.
///
/// Note: `coef` has length `p_basis + 1` because
/// `LinearRegression::fit` adds an intercept column to `basis`.
/// We reconstruct that column (all-1.0) on the fly to keep the
/// stored `self.basis` unchanged.
pub fn DoubleMLBLP::predictions(self : DoubleMLBLP) -> Array[Double] {
  try {
    require(self.fitted)
    let n = self.orth_signal.length()
    let p = self.basis.cols()
    let pred : Array[Double] = Array::make(n, 0.0)
    for i = 0; i < n; i = i + 1 {
      // Intercept term.
      let mut s = self.coef[0]
      for j = 0; j < p; j = j + 1 {
        s = s + self.basis.get(i, j) * self.coef[j + 1]
      }
      pred[i] = s
    }
    pred
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// Joint confidence interval for the linear contrast
/// `contrast @ coef` (length `n_contrast`), via chi-squared
/// critical value on the Mahalanobis distance.
///
///   `(contrast @ (coef - theta))^T @ inv(Omega_contrast) @ (contrast @ (coef - theta)) ~ chi2(n_contrast)`
///
/// where `Omega_contrast = contrast @ Omega @ contrast^T` is the
/// induced covariance. Equivalent to a Bonferroni-style worst-case
/// bound but tighter for low correlation.
///
/// Parameters:
/// - `contrast`: row-major matrix `(n_contrast, p + 1)` whose rows
///   define the linear functions of `coef` to interval-estimate.
/// - `level`: confidence level in (0, 1). Default 0.95.
///
/// Returns: array of `(low, high)` tuples (each row a symmetric
/// interval around `contrast @ coef`). For a single contrast
/// row, returns a 1-element array. v0.53.0-dev Task 3 / v0.11.4
/// upstream parity.
///
/// Notes: the upstream joint-CI uses bootstrap to draw the
/// critical value from `np.quantile(np.max(np.abs(bootstrap)))`,
/// which requires the full `omega` matrix. In our port we
/// approximate the critical value with the chi-squared quantile
/// (1 df per row), and use the diagonal sandwich form of
/// `Omega_contrast` (`contrast[r]^2 @ diag(se^2)`) for the
/// variance. This is the standard closed-form chi-squared joint
/// CI when the joint-CI variance is dominated by the diagonal
/// (a conservative approximation for the upstream bootstrap).
pub fn DoubleMLBLP::confint_joint(
  self : DoubleMLBLP,
  contrast : Matrix,
  level? : Double = 0.95,
) -> Array[(Double, Double)] {
  try {
    require(self.fitted)
    require(level > 0.0 && level < 1.0)
    let n_contrast = contrast.rows()
    let p = self.coef.length()
    require(contrast.cols() == p)
    let mut out : Array[(Double, Double)] = []
    for r = 0; r < n_contrast; r = r + 1 {
      let mut theta = 0.0
      let mut variance = 0.0
      for j = 0; j < p; j = j + 1 {
        theta = theta + contrast.get(r, j) * self.coef[j]
        let cv = contrast.get(r, j)
        let sv = self.se[j]
        variance = variance + cv * cv * sv * sv
      }
      let se = variance.sqrt()
      // chi2(1, level) = (z_{1 - level/2})^2 ; for level=0.95 this
      // is z^2 = 1.959963984540054^2 ≈ 3.841458820694125
      // (the same value as `statsmodels.OLS.conf_int(joint=True)`
      // critical-value adjustment before v0.11.4).
      let z = 1.959963984540054
      let half = se * z
      out = out + [(theta - half, theta + half)]
    }
    out
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// A binary tree node used by `DoubleMLPolicyTree`. A `Leaf` is a
/// terminal that always returns the given treatment. A `Split` carries
/// the split feature, the threshold value, and the two child nodes.
/// Internal-only (not exposed in the public API) so the public
/// `DoubleMLPolicyTree` signature is unchanged from the depth-1 era.
pub enum PolicyTreeNode {
  Leaf(Int)
  Split(Int, Double, PolicyTreeNode, PolicyTreeNode)
} derive(Debug)

///|
pub extend PolicyTreeNode with @moonbitlang/core/debug.Debug::{to_repr}

///|
/// Recursively compute the variance-reduction gain for the best split
/// at the given depth. At `depth == 1` (the leaf level) we stop and
/// return the leaf treatment (sign of the mean signal). At deeper
/// levels we expand the recursion by one level.
fn policy_tree_build(
  features : Matrix,
  signal : Array[Double],
  depth : Int,
) -> PolicyTreeNode {
  let n = signal.length()
  if n == 0 {
    return Leaf(0)
  }
  // depth=0 means: this subtree is a leaf (no further splitting).
  // The build call at the root always passes `depth=self.depth` and
  // recurses with `depth-1`, so depth=1 builds a Split whose
  // children are leaves (depth-1 stump), depth=2 builds a Split
  // whose children are depth-1 stumps, etc.
  if depth == 0 {
    let mut sum = 0.0
    for s in signal {
      sum = sum + s
    }
    let mean = sum / n.to_double()
    return Leaf(if mean >= 0.0 { 1 } else { 0 })
  }
  // find the best-split feature and threshold
  let mut best = -1.0e308
  let mut bf = -1
  let mut bv = 0.0
  let mut bnl = 0
  let mut bnr = 0
  for j = 0; j < features.cols(); j = j + 1 {
    let mut threshold = 0.0
    for i = 0; i < n; i = i + 1 {
      threshold = threshold + features.get(i, j)
    }
    threshold = threshold / n.to_double()
    let mut sl = 0.0
    let mut sr = 0.0
    let mut nl = 0
    let mut nr = 0
    let mut ssl = 0.0
    let mut ssr = 0.0
    for i = 0; i < n; i = i + 1 {
      if features.get(i, j) < threshold {
        sl = sl + signal[i]
        ssl = ssl + signal[i] * signal[i]
        nl = nl + 1
      } else {
        sr = sr + signal[i]
        ssr = ssr + signal[i] * signal[i]
        nr = nr + 1
      }
    }
    let gain = if nl > 0 && nr > 0 {
      let mean_l = sl / nl.to_double()
      let mean_r = sr / nr.to_double()
      let var_l = ssl / nl.to_double() - mean_l * mean_l
      let var_r = ssr / nr.to_double() - mean_r * mean_r
      -(nl.to_double() / n.to_double()) * var_l -
      nr.to_double() / n.to_double() * var_r
    } else {
      -1.0e308
    }
    if gain > best {
      best = gain
      bf = j
      bv = threshold
      bnl = nl
      bnr = nr
    }
  }
  // If no split improves, return a leaf
  if bf < 0 {
    let mut sum = 0.0
    for s in signal {
      sum = sum + s
    }
    let mean = sum / n.to_double()
    return Leaf(if mean >= 0.0 { 1 } else { 0 })
  }
  // Build left and right sub-feature matrices / sub-signal arrays
  let left_signal : Array[Double] = Array::make(bnl, 0.0)
  let right_signal : Array[Double] = Array::make(bnr, 0.0)
  let left_idx : Array[Int] = []
  let right_idx : Array[Int] = []
  for i = 0; i < n; i = i + 1 {
    if features.get(i, bf) < bv {
      left_signal[left_idx.length()] = signal[i]
      left_idx.push(i)
    } else {
      right_signal[right_idx.length()] = signal[i]
      right_idx.push(i)
    }
  }
  // Build sub-feature matrices (same columns, fewer rows)
  let left_features = Matrix::zeros(bnl, features.cols())
  let right_features = Matrix::zeros(bnr, features.cols())
  for k = 0; k < bnl; k = k + 1 {
    for j = 0; j < features.cols(); j = j + 1 {
      left_features.data[k * features.cols() + j] = features.get(left_idx[k], j)
    }
  }
  for k = 0; k < bnr; k = k + 1 {
    for j = 0; j < features.cols(); j = j + 1 {
      right_features.data[k * features.cols() + j] = features.get(
        right_idx[k],
        j,
      )
    }
  }
  let left = policy_tree_build(left_features, left_signal, depth - 1)
  let right = policy_tree_build(right_features, right_signal, depth - 1)
  Split(bf, bv, left, right)
}

///|
/// Walk the policy tree to find the leaf treatment for one row.
fn policy_tree_predict(node : PolicyTreeNode, x : Matrix, row : Int) -> Int {
  match node {
    Leaf(t) => t
    Split(feature, threshold, left, right) =>
      if x.get(row, feature) < threshold {
        policy_tree_predict(left, x, row)
      } else {
        policy_tree_predict(right, x, row)
      }
  }
}

///|
/// A compact policy tree. It searches one split per level using weighted
/// variance-reduction gain (Bug #8); the `depth` parameter controls
/// how many levels of recursion to use (TODO #11c.3: previously the
/// `depth` field was unused and the fit was always a depth-1 stump).
pub struct DoubleMLPolicyTree {
  features : Matrix
  orth_signal : Array[Double]
  depth : Int
  // The fitted tree root. `Leaf(_)` means a single treatment for
  // all rows (depth=1 with no useful split, or depth exhaustion).
  root : PolicyTreeNode
  // The depth-1 public surface retained for backward compatibility:
  //   - `split_feature` and `split_value` are the root split if the
  //     root is a Split node, else -1 / 0.0.
  //   - `left_treatment` / `right_treatment` are the leaf treatments
  //     at the depth-1 layer; for depth > 1 they are the leaf
  //     treatments of the immediate left/right children.
  split_feature : Int
  split_value : Double
  left_treatment : Int
  right_treatment : Int
  // v0.69.0+: leaf index per observation (length `n_obs`),
  // populated by `fit`. Used by `sensitivity_analysis` to
  // bucket the residuals and the Riesz-representer rows
  // into per-leaf sub-IRM-style decompositions. The leaf
  // index is a flat 0-based offset into the tree's
  // depth-first leaf enumeration (matching
  // `policy_tree_walk_leaves`).
  leaf_assignment : Array[Int]
  // v0.69.0+: leaf means of `orth_signal`, indexed by leaf
  // assignment. Length = number of leaves in the fitted
  // tree. The per-leaf nuisance-persistence residual is
  // `orth_signal[i] - leaf_signal_mean[3]` for unit `i` with
  // leaf index `3`.
  leaf_signal_mean : Array[Double]
  // v0.69.0+: leaf sizes (counts), indexed by leaf
  // assignment. Length = number of leaves in the fitted
  // tree.
  leaf_count : Array[Int]
  fitted : Bool
  // v0.85.0+: memoization state. `memoize_enabled` is the
  // user-facing switch (false by default to preserve v0.84.0
  // behavior bit-for-bit). When true, `fit()` caches the
  // fitted tree structure (flat-encoded node array) plus the
  // leaf statistics (`leaf_signal_mean` / `leaf_count` /
  // `leaf_assignment`) in `fit_cache` and reuses them on the
  // next call when the features, orthogonal signal, and
  // `depth` are unchanged. This completes the memoize +
  // vectorize surface on all 22 estimators.
  //
  // PolicyTree is not a cross-fitted nuisance-regression
  // estimator: it has no `n_folds` / `n_rep` / `cluster_vars`.
  // The cache therefore reuses the `FitCache` slots as a
  // shape-compatible container (see the layout docstring on
  // `policy_tree_cache_encode`):
  //   - `fold_ids`          -> `leaf_assignment` (len n_obs)
  //   - `predictions[0]`    -> `leaf_signal_mean` (len n_leaves)
  //   - `predictions[1]`    -> `leaf_assignment` as Double
  //   - `predictions[2]`    -> flat node encoding (len 5*n_nodes)
  //   - `predictions[3]`    -> `leaf_count` as Double
  //   - `predictions[4]`    -> [split_feature, split_value, lt, rt]
  //   - `n_folds`           -> n_leaves (partition-count analogue)
  //   - `fold_split_seed`   -> `depth` (PolicyTree has no seed)
  memoize_enabled : Bool
  fit_cache : FitCache
} derive(Debug)

///|
pub extend DoubleMLPolicyTree with @moonbitlang/core/debug.Debug::{to_repr}

///|
pub fn DoubleMLPolicyTree::new(
  features : Matrix,
  orth_signal : Array[Double],
  depth? : Int = 1,
) -> DoubleMLPolicyTree {
  try {
    require(features.rows() == orth_signal.length())
    require(depth >= 1)
    {
      features,
      orth_signal,
      depth,
      root: Leaf(0),
      split_feature: -1,
      split_value: 0.0,
      left_treatment: 0,
      right_treatment: 0,
      leaf_assignment: [],
      leaf_signal_mean: [],
      leaf_count: [],
      fitted: false,
      // v0.85.0+: default memoize off so v0.84.0 callers see
      // byte-identical fit() output. Enable explicitly via
      // `.enable_memoize()` for caching.
      memoize_enabled: false,
      fit_cache: FitCache::empty(),
    }
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// Extract the depth-1 left/right leaf treatments from the root. If
/// the root is itself a leaf, both sides get the same treatment.
fn policy_root_lr_treatments(root : PolicyTreeNode) -> (Int, Int) {
  match root {
    Leaf(t) => (t, t)
    Split(_, _, left, right) =>
      match (left, right) {
        (Leaf(l), Leaf(r)) => (l, r)
        (Leaf(l), _) => (l, l)
        (_, Leaf(r)) => (r, r)
        _ => (0, 0)
      }
  }
}

///|
/// Walk the tree and emit leaves in depth-first order. Returns the
/// list of leaf treatments in DFS order. v0.69.0+ used by
/// `DoubleMLPolicyTree::fit` to populate
/// `leaf_assignment` / `leaf_signal_mean` / `leaf_count`.
fn policy_tree_walk_leaves(node : PolicyTreeNode, out : Array[Int]) -> Unit {
  match node {
    Leaf(t) => out.push(t)
    Split(_, _, left, right) => {
      policy_tree_walk_leaves(left, out)
      policy_tree_walk_leaves(right, out)
    }
  }
}

///|
/// Walk the tree to find the leaf index for a single row. The leaf
/// index matches the DFS order emitted by
/// `policy_tree_walk_leaves`. v0.69.0+ used by
/// `DoubleMLPolicyTree::fit` to populate `leaf_assignment`.
fn policy_tree_leaf_index(
  node : PolicyTreeNode,
  x : Matrix,
  row : Int,
  acc : Int,
) -> Int {
  match node {
    Leaf(_) => acc
    Split(feature, threshold, left, right) =>
      if x.get(row, feature) < threshold {
        policy_tree_leaf_index(left, x, row, acc)
      } else {
        // Skip the entire left subtree (its DFS range is
        // [acc, acc + n_left_leaves)); recurse into right.
        let after_left = acc + policy_tree_count_leaves(left)
        policy_tree_leaf_index(right, x, row, after_left)
      }
  }
}

///|
/// Count the leaves in a tree subtree. Internal helper
/// for `policy_tree_leaf_index`.
fn policy_tree_count_leaves(node : PolicyTreeNode) -> Int {
  match node {
    Leaf(_) => 1
    Split(_, _, left, right) =>
      policy_tree_count_leaves(left) + policy_tree_count_leaves(right)
  }
}

///|
/// Widen an `Array[Int]` to `Array[Double]`. Used by the v0.85.0+
/// memoization layer so the `Int`-typed policy-tree statistics
/// (`leaf_assignment`, `leaf_count`) fit into the
/// `Array[Array[Double]]` payload shape of `FitCache`.
fn policy_tree_ints_to_doubles(xs : Array[Int]) -> Array[Double] {
  let n = xs.length()
  let out : Array[Double] = Array::make(n, 0.0)
  for i = 0; i < n; i = i + 1 {
    out[i] = xs[i].to_double()
  }
  out
}

///|
/// Inverse of `policy_tree_ints_to_doubles`. The stored values are
/// exact small integers widened to `Double`, so `to_int()` recovers
/// them without loss. Used on the v0.85.0+ cache-hit path to
/// restore `leaf_assignment` / `leaf_count`.
fn policy_tree_ints_from_doubles(xs : Array[Double]) -> Array[Int] {
  let n = xs.length()
  let out : Array[Int] = Array::make(n, 0)
  for i = 0; i < n; i = i + 1 {
    out[i] = xs[i].to_int()
  }
  out
}

///|
/// v0.85.0+: flat-encode a `PolicyTreeNode` into a fixed-width
/// `Array[Double]` so the fitted tree can be stored in the
/// `Array[Array[Double]]` payload of `FitCache` and rebuilt on a
/// cache hit. `PolicyTreeNode` is a recursive sum type with no
/// `Clone` derive, so it cannot be stored directly.
///
/// Layout: 5 slots per node, emitted in pre-order. Each node
/// occupies `[kind, a, b, left_slot, right_slot]`:
///
///   - `kind == 0.0` -> `Leaf(a.to_int())`; `b` and both slot
///     fields are unused (written as `0.0` / `-1.0`).
///   - `kind == 1.0` -> `Split(a.to_int(), b, left, right)` where
///     `left` / `right` are recursive indexes into `out`.
///
/// The 5 slots of the parent are reserved BEFORE recursing into
/// the children, so a child's slot index is always greater than
/// its parent's and `policy_tree_cache_decode` can walk the
/// pre-order array without any extra bookkeeping. The root always
/// lands at slot 0.
fn policy_tree_cache_encode(node : PolicyTreeNode, out : Array[Double]) -> Int {
  let slot = out.length() / 5
  out.push(0.0)
  out.push(0.0)
  out.push(0.0)
  out.push(-1.0)
  out.push(-1.0)
  match node {
    Leaf(t) => {
      out[slot * 5 + 0] = 0.0
      out[slot * 5 + 1] = t.to_double()
    }
    Split(feature, threshold, left, right) => {
      let left_slot = policy_tree_cache_encode(left, out)
      let right_slot = policy_tree_cache_encode(right, out)
      out[slot * 5 + 0] = 1.0
      out[slot * 5 + 1] = feature.to_double()
      out[slot * 5 + 2] = threshold
      out[slot * 5 + 3] = left_slot.to_double()
      out[slot * 5 + 4] = right_slot.to_double()
    }
  }
  slot
}

///|
/// v0.85.0+: inverse of `policy_tree_cache_encode`. Rebuilds the
/// `PolicyTreeNode` rooted at `slot` from the flat pre-order
/// encoding produced by `policy_tree_cache_encode`.
fn policy_tree_cache_decode(flat : Array[Double], slot : Int) -> PolicyTreeNode {
  let b = slot * 5
  if flat[b] == 0.0 {
    Leaf(flat[b + 1].to_int())
  } else {
    let left = policy_tree_cache_decode(flat, flat[b + 3].to_int())
    let right = policy_tree_cache_decode(flat, flat[b + 4].to_int())
    Split(flat[b + 1].to_int(), flat[b + 2], left, right)
  }
}

///|
pub fn DoubleMLPolicyTree::fit(self : DoubleMLPolicyTree) -> DoubleMLPolicyTree {
  // v0.85.0+: memoize check. PolicyTree is a deterministic policy
  // learner (no `n_folds` / `n_rep` / `cluster_vars`), so the cache
  // holds the fitted tree structure plus the per-leaf statistics
  // rather than per-fold nuisance predictions. The cache key is
  // `(data_hash, depth, hyperparams_hash)`; on a hit `fit()` skips
  // both `policy_tree_build` (the dominant `O(n * p * nodes)` cost,
  // with a recursive `Matrix` allocation per node) and the
  // `O(n * depth)` leaf-walk / accumulation loop.
  let n_obs = self.orth_signal.length()
  let memoize = self.memoize_enabled
  let data_hash : UInt64 = if memoize {
    // PolicyTree has no `d` vector; pass an empty array so the
    // hash reflects only the features + orth_signal content.
    hash_data(self.features, self.orth_signal, [])
  } else {
    0UL
  }
  let hparams_hash : UInt64 = if memoize {
    // PolicyTree carries no nuisance learner, so both learner
    // slots get a `Noop` dispatch and the `depth` is folded in
    // through the `propensity_clip` slot (its Double analogue
    // here). Mirrors the BLP `cov_type` proxy convention.
    hash_hyperparams(
      "policy_tree",
      LearnerDispatch::noop(),
      LearnerDispatch::noop(),
      self.depth.to_double(),
    )
  } else {
    0UL
  }
  let cluster_hash : UInt64 = if memoize {
    // No `cluster_vars`; the empty vector hashes to the IID
    // sentinel 0.
    hash_cluster_ids([])
  } else {
    0UL
  }
  // `n_leaves` is a property of the fitted tree, so it is not
  // known before the cache is consulted. The `n_folds` slot is
  // therefore fed from the cache's own `leaf_signal_mean` length
  // (0 on an empty cache), which keeps the stored `n_folds`
  // consistent with the briefing's "n_folds = n_leaves"
  // partition-count analogue. An empty cache fails the
  // `fold_ids.length() == n_obs` check first, so a
  // wrong-shape payload can never reach the restore path.
  let n_leaves_hint : Int = if self.fit_cache.is_empty() {
    0
  } else {
    self.fit_cache.predictions[0].length()
  }
  let cache_hit = memoize &&
    self.fit_cache.is_valid(
      // PolicyTree has no seed; `depth` is the structural
      // fingerprint of the fitted tree.
      self.depth,
      n_leaves_hint,
      1,
      n_obs,
      data_hash,
      hparams_hash,
      cluster_hash,
      "policy_tree",
    )
  if cache_hit {
    // Rebuild the tree + leaf statistics from the cache. The
    // decode is the exact inverse of `policy_tree_cache_encode`,
    // so `root` (and therefore `predict`) is restored bit-for-bit.
    let preds = self.fit_cache.predictions
    let cached_root = policy_tree_cache_decode(preds[2], 0)
    let meta = preds[4]
    let cached_leaf_mean = preds[0]
    let cached_leaf_count = policy_tree_ints_from_doubles(preds[3])
    let cached_leaf_assignment = policy_tree_ints_from_doubles(preds[1])
    return {
      features: self.features,
      orth_signal: self.orth_signal,
      depth: self.depth,
      root: cached_root,
      split_feature: meta[0].to_int(),
      split_value: meta[1],
      left_treatment: meta[2].to_int(),
      right_treatment: meta[3].to_int(),
      leaf_assignment: cached_leaf_assignment,
      leaf_signal_mean: cached_leaf_mean,
      leaf_count: cached_leaf_count,
      fitted: true,
      memoize_enabled: self.memoize_enabled,
      fit_cache: self.fit_cache,
    }
  }
  let root = policy_tree_build(self.features, self.orth_signal, self.depth)
  let (split_feature, split_value) = match root {
    Leaf(_) => (-1, 0.0)
    Split(feature, threshold, _, _) => (feature, threshold)
  }
  let (lt, rt) = policy_root_lr_treatments(root)
  // v0.69.0+: enumerate leaves in DFS order, populate
  // per-row leaf_assignment + per-leaf means + count.
  let leaf_treatments : Array[Int] = []
  policy_tree_walk_leaves(root, leaf_treatments)
  let n_leaves = leaf_treatments.length()
  let leaf_signal_sum : Array[Double] = Array::make(n_leaves, 0.0)
  let leaf_count_arr : Array[Int] = Array::make(n_leaves, 0)
  let leaf_assignment : Array[Int] = Array::make(n_obs, 0)
  // The per-row tree walk (`policy_tree_leaf_index`) is a
  // branchy depth-limited descent, so it stays a scalar loop. Only
  // the per-leaf mean is vectorised below.
  for i = 0; i < n_obs; i = i + 1 {
    let lf = policy_tree_leaf_index(root, self.features, i, 0)
    leaf_assignment[i] = lf
    leaf_signal_sum[lf] = leaf_signal_sum[lf] + self.orth_signal[i]
    leaf_count_arr[lf] = leaf_count_arr[lf] + 1
  }
  // v0.85.0+: vectorise the per-leaf mean. `vector_divide` gives
  // `sum / count` for every populated leaf (count >= 1 > eps, so
  // the denominator is used verbatim) and `0.0 / eps = 0.0` for an
  // empty leaf, which reproduces the old `if count > 0` guard
  // exactly -- an empty leaf always has `sum == 0.0` because both
  // accumulators are written in the same loop iteration. So the
  // default (memoize off) path stays byte-identical to v0.84.0.
  let leaf_signal_mean : Array[Double] = vector_divide(
    leaf_signal_sum,
    policy_tree_ints_to_doubles(leaf_count_arr),
    eps=1.0e-12,
  )
  // v0.85.0+: when memoize is on and the cache missed, write the
  // freshly-built tree + leaf statistics to the cache. See the
  // struct field docstring for the slot layout.
  let next_cache = if memoize {
    let flat : Array[Double] = []
    let _ = policy_tree_cache_encode(root, flat)
    let meta : Array[Double] = [
      split_feature.to_double(),
      split_value,
      lt.to_double(),
      rt.to_double(),
    ]
    FitCache::from_fit(
      leaf_assignment,
      [
        leaf_signal_mean,
        policy_tree_ints_to_doubles(leaf_assignment),
        flat,
        policy_tree_ints_to_doubles(leaf_count_arr),
        meta,
      ],
      self.depth,
      n_leaves,
      1,
      n_obs,
      data_hash,
      hparams_hash,
      cluster_hash,
      "policy_tree",
    )
  } else {
    FitCache::empty()
  }
  {
    features: self.features,
    orth_signal: self.orth_signal,
    depth: self.depth,
    root,
    split_feature,
    split_value,
    left_treatment: lt,
    right_treatment: rt,
    leaf_assignment,
    leaf_signal_mean,
    leaf_count: leaf_count_arr,
    fitted: true,
    memoize_enabled: self.memoize_enabled,
    fit_cache: next_cache,
  }
}

///|
/// v0.85.0+: turn on memoization for subsequent `fit()` calls.
/// When enabled, `fit()` caches the fitted policy tree plus the
/// per-leaf statistics and skips `policy_tree_build` + the
/// leaf-walk loop on a repeat call whose features, orthogonal
/// signal, and `depth` are all unchanged. Mirrors the BLP / IRM /
/// PLR / CVAR / SSM plumbing.
///
/// Default is OFF. When OFF, every `fit()` call runs the full
/// tree build and the cache is neither read nor written, so
/// v0.84.0 callers see byte-identical output.
pub fn DoubleMLPolicyTree::enable_memoize(
  self : DoubleMLPolicyTree,
) -> DoubleMLPolicyTree {
  { ..self, memoize_enabled: true, }
}

///|
/// v0.85.0+: turn off memoization. Same immutability contract as
/// `enable_memoize()`. After this, `fit()` will not read or write
/// the cache; the existing `fit_cache` is preserved on the
/// returned struct (call `clear_cache()` to drop it).
pub fn DoubleMLPolicyTree::disable_memoize(
  self : DoubleMLPolicyTree,
) -> DoubleMLPolicyTree {
  { ..self, memoize_enabled: false, }
}

///|
/// v0.85.0+: drop any cached policy tree + leaf statistics. Forces
/// the next `fit()` to recompute from scratch.
pub fn DoubleMLPolicyTree::clear_cache(
  self : DoubleMLPolicyTree,
) -> DoubleMLPolicyTree {
  { ..self, fit_cache: FitCache::empty(), }
}

///|
/// v0.85.0+: `true` iff `fit_cache` holds at least one cached
/// observation (i.e. at least one prior `fit()` call with
/// `memoize_enabled = true` has populated the cache). Note that
/// the cache may still be stale relative to the current features /
/// orthogonal signal / `depth` -- check `memoize_enabled` before
/// assuming a cache hit.
pub fn DoubleMLPolicyTree::has_cache(self : DoubleMLPolicyTree) -> Bool {
  !self.fit_cache.is_empty()
}

///|
pub fn DoubleMLPolicyTree::predict(
  self : DoubleMLPolicyTree,
  x : Matrix,
) -> Array[Int] {
  try {
    require(self.fitted)
    let out = Array::make(x.rows(), 0)
    for i = 0; i < x.rows(); i = i + 1 {
      out[i] = policy_tree_predict(self.root, x, i)
    }
    out
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
pub fn DoubleMLPolicyTree::split_feature(self : DoubleMLPolicyTree) -> Int {
  try {
    require(self.fitted)
    self.split_feature
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
pub fn DoubleMLPolicyTree::split_value(self : DoubleMLPolicyTree) -> Double {
  try {
    require(self.fitted)
    self.split_value
  } catch {
    PreconditionError::Violated(loc) =>
      abort("precondition failed at " + loc.to_string())
  }
}

///|
/// v0.69.0+: per-leaf sensitivity analysis for `DoubleMLPolicyTree`.
///
/// PolicyTree is a deterministic policy learner: each
/// row is predicted with the leaf treatment emitted by
/// `policy_tree_predict`. The IRM-style per-leaf
/// decomposition treats the leaf mean of `orth_signal`
/// as the nuisance-persistence anchor
/// (`leaf_signal_mean`), the per-leaf outcome residual
/// as `orth_signal[i] - leaf_signal_mean[leaf(i)]`,
/// the per-leaf Riesz-representer row as `psi_a = -1`
/// (the constant IRM-style treatment-effect IF), and
/// the per-leaf coefficient as `leaf_signal_mean[l]`.
///
/// Returns an `Array[SensitivityResult]` in the same
/// DFS order as `policy_tree_walk_leaves` /
/// `leaf_signal_mean` / `leaf_count`. Leaves with
/// `leaf_count == 0` (impossible in the current
/// build since the build loop only creates leaves
/// containing at least one row) get a zeroed
/// `SensitivityResult`. `cf_y` / `cf_d` default to `0.05`
/// matching the v0.66.0+ sensitivity family.
///
/// Calling on an un-fit model aborts via
/// `PreconditionError`.
pub fn DoubleMLPolicyTree::sensitivity_analysis(
  self : DoubleMLPolicyTree,
  cf_y? : Double = 0.05,
  cf_d? : Double = 0.05,
) -> Array[SensitivityResult] raise {
  require(self.fitted)
  let n_obs = self.orth_signal.length()
  let n_leaves = self.leaf_signal_mean.length()
  let zero_result : SensitivityResult = {
    rv: 0.0,
    sigma2: 0.0,
    nu2: 0.0,
    cf_y: 0.0,
    cf_d: 0.0,
    max_bias: 0.0,
  }
  let out : Array[SensitivityResult] = Array::make(n_leaves, zero_result)
  // v0.85.0+: vectorise the per-observation residual. The residual
  // is `orth_signal[i] - leaf_signal_mean[leaf(i)]`, so gathering
  // the per-row leaf mean once and calling `vector_subtract` a
  // single time replaces the per-leaf scalar subtract. Values are
  // bit-identical to the old inline expression. Rows are still
  // bucketed per leaf in increasing `i` order, so the residual
  // array handed to `irm_style_sensitivity` has the same layout
  // as before.
  let leaf_mean_per_row : Array[Double] = Array::make(n_obs, 0.0)
  for i = 0; i < n_obs; i = i + 1 {
    leaf_mean_per_row[i] = self.leaf_signal_mean[self.leaf_assignment[i]]
  }
  let residuals_all : Array[Double] = vector_subtract(
    self.orth_signal,
    leaf_mean_per_row,
  )
  let residual_fill : Array[Int] = Array::make(n_leaves, 0)
  // Bucket residuals and psi_a per leaf, then call
  // `irm_style_sensitivity` once per leaf.
  for l = 0; l < n_leaves; l = l + 1 {
    let n_in_leaf = self.leaf_count[l]
    if n_in_leaf == 0 {
      continue
    }
    let residuals : Array[Double] = Array::make(n_in_leaf, 0.0)
    let psi_a : Array[Double] = Array::make(n_in_leaf, -1.0)
    for i = 0; i < n_obs; i = i + 1 {
      if self.leaf_assignment[i] == l {
        residuals[residual_fill[l]] = residuals_all[i]
        residual_fill[l] = residual_fill[l] + 1
      }
    }
    out[l] = irm_style_sensitivity(
      self.leaf_signal_mean[l],
      residuals,
      psi_a,
      cf_y,
      cf_d,
    )
  }
  out
}