metacausal.aggregation.Select¶
- class metacausal.aggregation.Select(split=<factory>, propensity_model=None, outcome_model=None, propensity_trim=0.01, fit_nuisance_fn=None, loss='dr', pseudo_outcome_fn=None)[source]¶
Bases:
SupervisedStrategySelect the single component model minimizing a pseudo-outcome risk.
Not an ensemble — a feasible selection procedure expressed as a
SupervisedStrategyso it shares the same fit-time contract (cross-fitted nuisances, out-of-fold CATE predictions) and downstream machinery (bootstrap,EnsembleWeightsintrospection) as the true aggregators. It is the natural comparator for “does aggregation beat feasible selection?” experiments.Two risk criteria:
loss="dr"DR/AIPW plug-in risk,
mean_i (Gamma_i - tau_hat_k(x_i))**2. Atnu=1,QAggregation’s objective is exactly the weighted average of these per-model MSEs (see its docstring), which is linear in the simplex weights — so the constrained optimum sits at the one-hot vertex minimizing individual MSE.Select(loss="dr")computes that vertex directly by argmin rather than relying on a QP solver to find it, and is in the lineage of influence-function- corrected plug-in validation (Alaa & van der Schaar, 2019).loss="r"R-risk,
mean_i (Y_tilde_i - tau_hat_k(x_i)*W_tilde_i)**2, using the Robinson decomposition residuals (Nie & Wager, 2021; the selection criterion advocated by Schuler et al., 2018 and Doutreligne & Varoquaux, 2025).
- Parameters:
split (CrossFitSplit or TrainAvgSplit) – Data splitting strategy. Default: 5-fold cross-fitting.
propensity_model (sklearn classifier or None) – Model for P(T=1|X).
NoneselectsHistGradientBoostingClassifier.outcome_model (sklearn regressor / classifier or None) – Model for E[Y|X, T]. For continuous Y, must be a regressor; for binary Y, must be a classifier (
predict_proba).Noneselects an outcome-type-appropriate default (HistGradientBoostingRegressororHistGradientBoostingClassifier).propensity_trim (float) – Clip propensity scores to [trim, 1-trim].
fit_nuisance_fn (callable or None) – Optional replacement for the default nuisance fitting procedure. See
SupervisedStrategy.fit_nuisance_fnfor the required signature.loss (str) – Risk criterion:
"dr"(default) or"r".pseudo_outcome_fn (callable or None) –
Optional replacement for
dr_pseudo_outcome, used only whenloss="dr". Must have signature:pseudo_outcome_fn(Y, T, nuisance) -> ndarray of shape (m,)
None(default) uses the standard AIPW form.
Examples
>>> from sklearn.linear_model import LinearRegression >>> from sklearn.ensemble import HistGradientBoostingRegressor as HGBR >>> from metacausal import CausalEnsemble >>> from metacausal.adapters import GenericCATEAdapter >>> from metacausal.aggregation import Select >>> from metacausal.datasets import load_lalonde >>> X, T, Y = load_lalonde() >>> def fit_linear(X, T, Y, **kwargs): ... treated = T == 1 ... m1 = LinearRegression().fit(X[treated], Y[treated]) ... m0 = LinearRegression().fit(X[~treated], Y[~treated]) ... return (m1, m0) >>> def fit_hgb(X, T, Y, **kwargs): ... treated = T == 1 ... m1 = HGBR(max_iter=20).fit(X[treated], Y[treated]) ... m0 = HGBR(max_iter=20).fit(X[~treated], Y[~treated]) ... return (m1, m0) >>> def cate_fn(state, X): ... m1, m0 = state ... return m1.predict(X) - m0.predict(X) >>> methods = [ ... GenericCATEAdapter(fit_linear, cate_fn, name="linear"), ... GenericCATEAdapter(fit_hgb, cate_fn, name="hgb"), ... ] >>> ens = CausalEnsemble(methods=methods, aggregation=Select(loss="dr")) >>> _ = ens.fit(X, T, Y, random_state=42) >>> ens.ate() AteEstimate(ate=..., n_methods=2, aggregation='Select', spread=...)
Methods
__init__aggregateAggregate component CATEs using stored weights.
Select the lowest-risk component from OOF predictions and nuisance estimates.
Attributes
ensemble_weightsEnsemble weights computed during
fit().fit_nuisance_fnoutcome_modelpropensity_modelpropensity_trimstrategy_familysplitDetails
- fit_weights(cate_predictions, Y, T, X, nuisance)[source]¶
Select the lowest-risk component from OOF predictions and nuisance estimates.
- Parameters:
cate_predictions (array of shape (K, m))
Y (arrays of shape (m,))
T (arrays of shape (m,))
X (array of shape (m, p))
nuisance (NuisanceEstimates)
- Returns:
EnsembleWeights with a one-hot
weightsvector at the selectedmodel (model_names populated by
_fit_supervised);detailscarries the full per-model risk vector.
- Return type: