metacausal.plots.disagreement¶
- metacausal.plots.disagreement(ensemble, X, *, ax=None, metric='spearman', cluster=False, annotate=True)[source]¶
Pairwise disagreement between component CATEs evaluated on
X.Computes each CATE-capable component’s CATE on
X, forms a(K, K)matrix of pairwise agreement undermetric, and renders it as a heatmap with optional cell annotations.- Parameters:
ensemble (CausalEnsemble) – A fitted
CausalEnsemblewith at least two CATE-capable components.X (np.ndarray) – Covariates to evaluate each component’s CATE on, shape
(n, p). Typically the training data or a held-out sample.ax (Axes | None) – Existing axes to draw on. If
None, a new figure is created.metric (Literal['spearman', 'pearson', 'rmse']) –
Pairwise metric:
"spearman"(default): rank correlation of unit-level CATEs. Robust to scale differences between components."pearson": linear correlation of unit-level CATEs."rmse": root mean squared difference. Scale-aware and dominated by components with extreme predictions.
cluster (bool) – If
True, reorder rows/columns by hierarchical clustering. Correlation metrics use1 - |corr|as the distance; RMSE is used directly. Requiresscipy.cluster.hierarchy.annotate (bool) – If
True, write each cell’s value inside the heatmap.
- Returns:
Axes – The axes the plot was drawn on.
- Raises:
ValueError – If
ensemblehas fewer than two CATE-capable components.- Return type:
Axes
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.datasets import load_lalonde >>> from metacausal.plots import disagreement >>> 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) >>> _ = ens.fit(X, T, Y, random_state=42) >>> ax = disagreement(ens, X) >>> ax.get_title() 'Component CATE spearman agreement'
(
Source code,png)