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 under metric, and renders it as a heatmap with optional cell annotations.

Parameters:
  • ensemble (CausalEnsemble) – A fitted CausalEnsemble with 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 use 1 - |corr| as the distance; RMSE is used directly. Requires scipy.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 ensemble has 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)

../_images/metacausal-plots-disagreement-1.png