Causal GAM for treatment effect estimation.
CausalGAM(
outcome,
treatment,
confounders,
*,
method="partially_linear",
family=None,
n_folds=5,
)
Estimates the causal effect of a treatment D on an outcome Y, controlling for a set of confounders X, using double/debiased machine learning (DML; Chernozhukov et al. 2018) with GAM nuisance models. Two structural forms are supported:
- Partially linear (
method="partially_linear"): Y = theta * D + f(X) + eps, giving a single constant average treatment effect (ATE) theta.
- Interactive (
method="interactive"): Y = g(D, X) + eps, allowing the treatment effect to vary smoothly with X (conditional average treatment effect, CATE).
DML addresses the regularization bias that arises when flexible, penalized nuisance models (here, GAMs for E[Y | X] and E[D | X]) are plugged directly into a naive treatment-effect estimator: the smoothing bias in the nuisance fits would otherwise leak into the treatment effect estimate. Cross-fitting (fitting nuisance models on one subset of folds and evaluating residuals on the held-out fold) together with a Neyman-orthogonal moment condition makes the resulting ATE estimate root-n consistent and asymptotically normal even though the nuisance GAMs converge at slower nonparametric rates.
Use CausalGAM for observational-data effect estimation where confounding is plausibly captured by smooth functions of observed covariates, and you want valid inference (standard errors, confidence intervals) on the treatment effect rather than just a point prediction.
Parameters
outcome: str
-
Name of the outcome variable.
treatment: str
-
Name of the treatment variable.
confounders: list[str]
-
List of confounder variable names. Both the outcome and treatment nuisance GAMs use s(c) smooth terms for each confounder c.
method: str = "partially_linear"
-
"partially_linear" (default) for constant ATE, or "interactive" for heterogeneous treatment effects (enables .cate()).
family: Family | None = None
-
Response distribution for the outcome nuisance model. Defaults to Gaussian(). The treatment nuisance model always uses Gaussian() regardless of this setting, since DML residualizes the treatment via its conditional mean.
n_folds: int = 5
-
Number of cross-fitting folds for DML (default
5). Each fold’s nuisance models are fit on the other n_folds - 1 folds and evaluated on the held-out fold to avoid overfitting bias.
Notes
Fitting proceeds in three steps. First, cross-fitted residuals are formed for both outcome and treatment:
\hat\varepsilon_{Y,i} = Y_i - \hat m_Y(X_i), \qquad
\hat\varepsilon_{D,i} = D_i - \hat m_D(X_i)
where \hat m_Y and \hat m_D are GAM estimates of E[Y \mid X] and E[D \mid X], each fit on folds excluding observation i. Second, the ATE is estimated by the residual-on-residual regression (the partialling-out estimator):
\hat\theta = \frac{\sum_i \hat\varepsilon_{D,i} \, \hat\varepsilon_{Y,i}}
{\sum_i \hat\varepsilon_{D,i}^2}
Third, its standard error is derived from the empirical variance of the Neyman-orthogonal score \psi_i = \hat\varepsilon_{D,i}(\hat\varepsilon_{Y,i} - \hat\theta \hat\varepsilon_{D,i}):
\widehat{\mathrm{se}}(\hat\theta) = \sqrt{\frac{\overline{\psi^2}}
{\left(\sum_i \hat\varepsilon_{D,i}^2\right)^{2} / n}}
When method="interactive", a further GAM is fit on the pseudo-outcome \hat\varepsilon_{Y,i} / \hat\varepsilon_{D,i}, weighted by \hat\varepsilon_{D,i}^2, to recover the CATE as a smooth function of the confounders.
Examples
import numpy as np
from whittaker.causal import CausalGAM
rng = np.random.default_rng(0)
n = 1000
x = rng.uniform(0, 1, n)
d = rng.binomial(1, 1 / (1 + np.exp(-(2 * x - 1))), n).astype(float)
y = 1.5 * d + np.sin(2 * np.pi * x) + rng.normal(scale=0.3, size=n)
model = CausalGAM(outcome="y", treatment="d", confounders=["x"], n_folds=5)
model.fit({"x": x, "d": d, "y": y}, seed=0)
print(model.treatment_effect())
TreatmentEffect(ate=1.4973, se=0.0095, p=0.0000, 95% CI=[1.4786, 1.5160])
Attributes
|
Name
|
Description
|
|
confounders
|
Names of the confounder variables X controlled for in the nuisance GAMs.
|
|
is_fitted
|
Whether fit() has been called successfully on this model.
|
|
method
|
Structural form used for estimation.
|
|
outcome
|
Name of the outcome variable Y used when constructing this model.
|
|
treatment
|
Name of the treatment variable D used when constructing this model.
|
confounders
Names of the confounder variables X controlled for in the nuisance GAMs.
Returns a copy, so mutating the returned list does not affect the model.
is_fitted
Whether fit() has been called successfully on this model.
Most other methods (treatment_effect(), cate(), residuals(), summary()) raise RuntimeError if this is False.
method
Structural form used for estimation.
Either "partially_linear" (constant ATE) or "interactive" (heterogeneous treatment effects, enabling .cate()).
outcome
Name of the outcome variable Y used when constructing this model.
treatment
Name of the treatment variable D used when constructing this model.
Methods
|
Name
|
Description
|
|
cate()
|
Estimate conditional average treatment effects.
|
|
fit()
|
Fit the causal GAM via cross-fitted DML.
|
|
residuals()
|
Return the orthogonalized residuals.
|
|
summary()
|
Build a human-readable text summary of the fitted causal GAM.
|
|
treatment_effect()
|
Compute the average treatment effect with inference.
|
cate()
Estimate conditional average treatment effects.
cate(
new_data=None,
*,
variable=None,
n_points=100,
level=0.95,
)
Requires method="interactive". Evaluates the fitted CATE model (a GAM regressed on the pseudo-outcome \hat\varepsilon_Y / \hat\varepsilon_D) either on user-supplied new_data or on a grid over one confounder, holding the other confounders at their training-data means. Returns CATE as a function of a chosen confounder variable, together with pointwise confidence bands derived from the CATE model’s own standard errors.
Parameters
new_data: InputData | None = None
-
Covariate data for prediction. If None, evaluates on a grid of the specified variable.
variable: str | None = None
-
Confounder to condition on. Required if new_data is None. Defaults to the first confounder.
n_points: int = 100
-
Number of grid points (used when new_data is None).
level: float = 0.95
-
Confidence level for intervals.
fit()
Fit the causal GAM via cross-fitted DML.
fit(
data,
*,
fit_method="REML",
select=False,
seed=None,
)
Randomly assigns observations to n_folds folds, then for each fold fits an outcome GAM E[Y | X] and a treatment GAM E[D | X] on the remaining folds and predicts residuals on the held-out fold. The pooled cross-fitted residuals are combined into the partialling-out ATE estimate and its standard error (see class Notes). If method="interactive", an additional CATE model is fit on the residual ratio.
Parameters
data: InputData
-
Column-oriented data containing outcome, treatment, and confounder columns.
fit_method: str = "REML"
-
Smoothing parameter selection method for the outcome and treatment nuisance GAMs (default "REML").
select: bool = False
-
Enable double-penalty variable selection in the nuisance GAMs.
seed: int | None = None
-
Random seed for fold assignment.
Returns
CausalGAM
-
Returns
self for method chaining.
residuals()
Return the orthogonalized residuals.
Returns
tuple[NDArray, NDArray]
-
(residuals_y, residuals_d), a residualized outcome and treatment after partialling out confounders.
summary()
Build a human-readable text summary of the fitted causal GAM.
Reports the outcome, treatment, confounders, estimation method, number of cross-fitting folds, and the treatment effect (ATE, standard error, confidence interval, and p-value) from treatment_effect(). If method="interactive" and a CATE model was successfully fit, notes that .cate() can be used for estimates.
Returns
str
-
Multi-line summary text.
treatment_effect()
Compute the average treatment effect with inference.
treatment_effect(level=0.95)
Builds a TreatmentEffect from the ATE and standard error computed during fit(), adding a normal-approximation confidence interval and two-sided Wald p-value for H0: ATE = 0.
Parameters
level: float = 0.95
-
Confidence level (default
0.95).