library(hal9001)
library(glmnet)
library(furrr)
library(plotly)
library(origami)
library(dplyr)
library(tidyr)
library(ggplot2)Conditional treatment effects estimation with the HAL-based R-learner
Overview
Average treatment effects summarize how a treatment works across an entire population. In many applications, however, the treatment effect varies with a person’s baseline characteristics. The conditional average treatment effect (CATE) describes that heterogeneity.
In this tutorial, we provide an example of estimating the CATE with the R-learner. Here, highly adaptive lasso (HAL) basis functions provide a flexible representation of the covariates, and glmnet performs the Lasso fits.
By the end of this tutorial, you will be able to:
- construct a zero-order HAL design matrix;
- estimate the treatment propensity and marginal outcome regression;
- derive the weighted pseudo-outcome form of the R-loss;
- fit an R-learner with cross-validated Lasso regression; and
- evaluate the CATE.
1. Setup
We first load the packages we need for this tutorial. We will use hal9001 to enumerate and evaluate HAL basis functions, glmnet for cross-validated Lasso regression, origami for folds, and plotly for interactive surfaces.
If needed, install the packages before proceeding the tutorial.
install.packages(c(
"hal9001", "glmnet", "furrr", "plotly", "origami",
"dplyr", "tidyr", "ggplot2"
))2. Data-generating process and target parameter
We observe \(O=(W,A,Y)\), where \(W=(W_1,W_2,W_3)\) contains baseline covariates, \(A\in\{0,1\}\) is treatment, and \(Y\) is a continuous outcome. The covariates are independently generated on a grid from \(-1\) to \(1\). Specifically, \(W_j=\operatorname{round}(U_j,1)\), \(U_j\sim\operatorname{Uniform}(-1,1)\), for \(j=1,2,3.\)
Treatment follows \(A\mid W \sim \operatorname{Bernoulli}\{e_0(W)\}\) with \(e_0(W)=\operatorname{expit}(-0.25W_1+W_2).\)
The outcome is \(Y=-0.5+W_1+0.5W_2+0.3W_3+A\tau_0(W_1,W_2)+U_Y\) with \(U_Y\sim N(0,1),\) and the true CATE being \(\tau_0(W_1,W_2)=0.5W_1+\sin(2\pi W_2).\)
The effect changes linearly with \(W_1\) and periodically with \(W_2\). Although \(W_3\) predicts the outcome, it does not modify the treatment effect.
Interpreting the conditional contrast as a causal effect requires consistency, conditional exchangeability \((Y^0,Y^1)\perp A\mid W\), and positivity \(0<P(A=1\mid W=w)<1\) in the covariate region of interest. These conditions hold by construction in the simulation.
Simulate one data set
The helper file sim_data.R contains the complete data-generating process.
source("sim_data.R")
set.seed(1234)
n <- 2000
data <- sim_data(n)
W <- data[, grep("W", colnames(data))]
A <- data$A
Y <- data$YLet’s inspect the structure and basic summaries first.
dplyr::glimpse(data)Rows: 2,000
Columns: 5
$ W1 <dbl> -0.7, 0.4, -0.1, 0.3, -0.1, 0.9, 0.8, 0.6, -0.9, 0.3, -0.7, -0.2, -…
$ W2 <dbl> -0.3, -0.4, 0.4, -0.5, 0.2, 0.1, 0.9, 0.8, 0.9, 0.7, -0.3, 0.4, -0.…
$ W3 <dbl> -0.7, -0.1, -0.6, 0.5, -0.2, 0.6, 0.9, 0.7, 0.5, -0.5, 0.7, -0.3, -…
$ A <int> 0, 0, 0, 1, 1, 0, 1, 1, 1, 1, 1, 0, 1, 0, 1, 1, 1, 1, 0, 1, 1, 0, 1…
$ Y <dbl> -2.76706575, -0.05257076, 0.50444118, -2.49569770, 0.77018121, 1.13…
data.frame(
sample_size = nrow(data),
treatment_prevalence = mean(A),
outcome_mean = mean(Y),
outcome_sd = sd(Y)
) sample_size treatment_prevalence outcome_mean outcome_sd
1 2000 0.4755 -0.5381055 1.31629
The following exploratory plot shows the observed outcome against \(W_2\) by treatment group.
ggplot(data, aes(x = W2, y = Y, color = factor(A))) +
geom_jitter(width = 0.025, height = 0, alpha = 0.35, size = 1) +
geom_smooth(
aes(group = factor(A)),
method = "loess",
se = FALSE,
linewidth = 0.9
) +
scale_color_manual(
values = c("0" = "#2b8cbe", "1" = "#e34a33"),
labels = c("0" = "Untreated", "1" = "Treated")
) +
labs(
x = expression(W[2]),
y = "Observed outcome Y",
color = "Treatment",
title = "Raw outcome patterns do not isolate treatment effects"
) +
theme_minimal(base_size = 12) +
theme(legend.position = "bottom")3. Create cross-validation folds
All three primary Lasso fits use the same five-fold split.
# make folds
folds <- make_folds(n = n, V = 5)
foldid <- folds2foldvec(folds)Check that every observation belongs to one fold and that fold sizes are reasonably balanced.
fold_counts <- as.data.frame(table(foldid))
names(fold_counts) <- c("fold", "observations")
fold_counts fold observations
1 1 400
2 2 400
3 3 400
4 4 400
5 5 400
stopifnot(length(foldid) == n, !anyNA(foldid))Here, cv.glmnet() uses held-out folds to choose each Lasso penalty, but the final nuisance models are refit on the full sample and their fitted values are predicted on that same sample. A fully cross-fitted R-learner would instead use out-of-fold nuisance predictions for every observation. For pedagogical purposes, we go with a simpler implementation.
4. Construct the HAL basis
The first HAL basis represents functions of all three covariates. With smoothness_orders = 0L, the basis consists of indicator-type tensor-product functions. max_degree = 2L restricts covariate interactions up to two-way.
basis_list <- enumerate_basis(x = W,
max_degree = 2L,
smoothness_orders = 0L)
phi_W <- make_design_matrix(X = as.matrix(W),
blist = basis_list)Each row of phi_W corresponds to an observation. Each column is a candidate HAL basis function that glmnet may retain or shrink to zero.
data.frame(
observations = nrow(phi_W),
candidate_basis_functions = ncol(phi_W),
maximum_interaction_degree = 2L,
smoothness_order = 0L
) observations candidate_basis_functions maximum_interaction_degree
1 2000 1357 2
smoothness_order
1 0
5. Estimate the treatment propensity
Define \(e_0(W)=E_0(A\mid W)=P_0(A=1\mid W).\)
We estimate \(e_0\) with binomial Lasso regression on the HAL basis.
# estimate E(A|W)
g_fit <- cv.glmnet(x = phi_W,
y = A,
family = "binomial",
alpha = 1,
foldid = foldid)
g1W <- as.numeric(predict(g_fit, newx = phi_W, s = "lambda.min", type = "response"))g1W is the fitted probability of treatment for each observation.
The simulated sample has comfortable empirical overlap.
data.frame(
propensity = g1W,
treatment = factor(A, labels = c("Untreated", "Treated"))
) |>
ggplot(aes(x = propensity, fill = treatment)) +
geom_histogram(bins = 30, alpha = 0.65, position = "identity") +
facet_wrap(vars(treatment), nrow = 1) +
scale_fill_manual(values = c("Untreated" = "#2b8cbe", "Treated" = "#e34a33")) +
labs(
x = expression(hat(e)(W)),
y = "Observations",
fill = NULL,
title = "Overlap diagnostic"
) +
theme_minimal(base_size = 12) +
theme(legend.position = "none")6. Estimate the marginal outcome regression
The second nuisance function is \(m_0(W)=E_0(Y\mid W).\)
This is not the treatment-specific outcome regression \(E(Y\mid A,W)\). It is the outcome mean after averaging over the observed treatment mechanism. Under this simulation, \(m_0(W)=-0.5+W_1+0.5W_2+0.3W_3+e_0(W)\tau_0(W_1,W_2).\)
We use Gaussian Lasso regression on the same HAL design matrix.
# estimate E(Y|W)
theta_fit <- cv.glmnet(x = phi_W,
y = Y,
family = "gaussian",
alpha = 1,
foldid = foldid)
theta <- as.numeric(predict(theta_fit, newx = phi_W, s = "lambda.min", type = "response"))7. Derive and fit the R-loss
The Robinson decomposition motivates the R-learner: \(Y-m_0(W)=\{A-e_0(W)\}\tau_0(W)+\varepsilon,\) where the remaining error is conditionally mean zero under the partially linear treatment-effect model. Replacing the nuisance functions with estimates leads to the empirical R-loss
\[\mathcal L_R(\tau)=\sum_{i=1}^n\bigg\{Y_i-\hat m(W_i)-(A_i-\hat e(W_i))\tau(W_i)\bigg\}^2.\]
Weighted pseudo-outcome form
For observations with \(A_i\neq\hat e(W_i)\), define \(\widetilde Y_i=(Y_i-\hat m(W_i))/(A_i-\hat e(W_i))\) and \(\omega_i=\{A_i-\hat e(W_i)\}^2.\)
Then
\[\omega_i(\widetilde Y_i-\tau(W_i))^2=\bigg\{Y_i-\hat m(W_i)-\{A_i-\hat e(W_i)\}\tau(W_i)\bigg\}^2.\]
Thus, a weighted regression of the pseudo-outcome on \(W\) minimizes the R-loss.
# R-loss
pseudo_outcome <- (Y-theta)/(A-g1W)
pseudo_weights <- (A-g1W)^2We now fit a Gaussian Lasso of the pseudo-outcome on the full HAL basis using the R-loss weights.
r_lrnr_fit <- cv.glmnet(x = phi_W,
y = pseudo_outcome,
weights = pseudo_weights,
family = "gaussian",
alpha = 1,
foldid = foldid)
cate <- as.numeric(predict(r_lrnr_fit, newx = phi_W, s = "lambda.min", type = "response"))Assess fitted CATE values at the observed covariates
In a simulation, the true effect is available at every observed \(W\).
true_cate_observed <- 0.5 * data$W1 + sin(2 * pi * data$W2)
observed_cate_error <- data.frame(
MAE = mean(abs(cate - true_cate_observed)),
RMSE = sqrt(mean((cate - true_cate_observed)^2)),
correlation = cor(cate, true_cate_observed)
)
summary(cate) Min. 1st Qu. Median Mean 3rd Qu. Max.
-1.402539 -0.512010 -0.060810 -0.002658 0.584456 2.010944
observed_cate_error MAE RMSE correlation
1 0.2735193 0.344946 0.8881704
8. Summarize the CATE over effect modifiers
The primary R-learner used \(W_1\), \(W_2\), and \(W_3\). For visualization, we could fit a second HAL that summarizes cate as a function of \(W_1\) and \(W_2\) only.
# CATE on W1, W2
basis_list_W1W2 <- enumerate_basis(x = data[, c("W1", "W2")],
max_degree = 2L,
smoothness_orders = 0L)
phi_W1W2 <- make_design_matrix(X = as.matrix(data[, c("W1", "W2")]),
blist = basis_list_W1W2)
cate_W1W2_fit <- cv.glmnet(x = phi_W1W2,
y = cate,
family = "gaussian",
alpha = 1,
foldid = foldid)9. Evaluate the estimated CATE on a grid
We evaluate both the estimated and true treatment effects on all combinations of 21 values of \(W_1\) and 21 values of \(W_2\). This produces a \(21\times21=441\) point grid over the support used by the data generator.
W1_vals <- seq(-1, 1, 0.1)
W2_vals <- seq(-1, 1, 0.1)
grid <- expand.grid(W1 = W1_vals,
W2 = W2_vals)
grid_hal_design <- make_design_matrix(X = as.matrix(grid),
blist = basis_list_W1W2)
grid$pred_cate <- as.numeric(predict(cate_W1W2_fit, newx = grid_hal_design,
s = "lambda.min", type = "response"))
grid$true_cate <- sim_data(n, grid = grid[, c("W1", "W2")])Compare one-dimensional slices
Holding \(W_1\) fixed makes the sinusoidal dependence on \(W_2\) easy to see.
slice_W1_vals <- c(-1, 0, 1)
grid |>
filter(W1 %in% slice_W1_vals) |>
select(W1, W2, pred_cate, true_cate) |>
pivot_longer(
cols = c(pred_cate, true_cate),
names_to = "curve",
values_to = "CATE"
) |>
mutate(
curve = recode(
curve,
pred_cate = "Estimated CATE",
true_cate = "True CATE"
),
W1_slice = factor(
sprintf("W1 = %.1f", W1),
levels = sprintf("W1 = %.1f", slice_W1_vals)
)
) |>
ggplot(aes(x = W2, y = CATE, color = curve, linetype = curve)) +
geom_hline(yintercept = 0, color = "grey75", linewidth = 0.5) +
geom_line(linewidth = 0.9) +
facet_wrap(vars(W1_slice), nrow = 1) +
scale_color_manual(values = c(
"Estimated CATE" = "#238b45",
"True CATE" = "#cb181d"
)) +
labs(
x = expression(W[2]),
y = "Conditional average treatment effect",
color = NULL,
linetype = NULL
) +
theme_minimal(base_size = 12) +
theme(legend.position = "bottom")The true curves have the same sinusoidal shape and are vertically separated by \(0.5W_1\).
10. Build interactive CATE surfaces
pred_cate_mat <- matrix(grid$pred_cate,
nrow = length(W2_vals),
ncol = length(W1_vals),
byrow = TRUE)
true_cate_mat <- matrix(grid$true_cate,
nrow = length(W2_vals),
ncol = length(W1_vals),
byrow = TRUE)
stopifnot(
identical(dim(pred_cate_mat), c(length(W2_vals), length(W1_vals))),
identical(dim(true_cate_mat), c(length(W2_vals), length(W1_vals)))
)plt_pred <- plot_ly(x = W1_vals, y = W2_vals, z = pred_cate_mat,
colorscale = "Greens") %>%
add_surface() %>%
layout(title = list(text = "Estimated CATE",
font = list(size = 18, color = "black"),
yref = "container",
y = 0.98),
scene = list(xaxis = list(title = "W1", titlefont = list(size = 14)),
yaxis = list(title = "W2", titlefont = list(size = 14)),
zaxis = list(title = "CATE", titlefont = list(size = 14), range = c(-1.5, 1.5)),
camera = list(eye = list(x = -1.5, y = 1.5, z = 0.8))),
legend = list(x = 0.8, y = 1, font = list(size = 12)))
plt_true <- plot_ly(x = W1_vals, y = W2_vals, z = true_cate_mat,
colorscale = "Reds") %>%
add_surface() %>%
layout(title = list(text = "True CATE",
font = list(size = 18, color = "black"),
yref = "container",
y = 0.98),
scene = list(xaxis = list(title = "W1", titlefont = list(size = 14)),
yaxis = list(title = "W2", titlefont = list(size = 14)),
zaxis = list(title = "CATE", titlefont = list(size = 14), range = c(-1.5, 1.5)),
camera = list(eye = list(x = -1.5, y = 1.5, z = 0.8))),
legend = list(x = 0.8, y = 1, font = list(size = 12)))Estimated CATE surface
Drag to rotate, scroll to zoom, and hover to read individual values.
plt_predTrue CATE surface
The true surface combines the plane \(0.5W_1\) with the wave \(\sin(2\pi W_2)\).
plt_trueReproducibility information
Code
sessionInfo()R version 4.5.2 (2025-10-31)
Platform: aarch64-apple-darwin20
Running under: macOS Tahoe 26.2
Matrix products: default
BLAS: /System/Library/Frameworks/Accelerate.framework/Versions/A/Frameworks/vecLib.framework/Versions/A/libBLAS.dylib
LAPACK: /Library/Frameworks/R.framework/Versions/4.5-arm64/Resources/lib/libRlapack.dylib; LAPACK version 3.12.1
locale:
[1] C.UTF-8/C.UTF-8/C.UTF-8/C/C.UTF-8/C.UTF-8
time zone: America/Los_Angeles
tzcode source: internal
attached base packages:
[1] stats graphics grDevices utils datasets methods base
other attached packages:
[1] tidyr_1.3.2 dplyr_1.2.1 origami_1.0.7 plotly_4.12.1 ggplot2_4.0.3
[6] furrr_0.3.1 future_1.68.0 glmnet_5.0 Matrix_1.7-4 hal9001_0.4.6
[11] Rcpp_1.1.1
loaded via a namespace (and not attached):
[1] generics_0.1.4 shape_1.4.6.1 stringi_1.8.7
[4] lattice_0.22-7 listenv_0.10.0 digest_0.6.39
[7] magrittr_2.0.4 evaluate_1.0.5 grid_4.5.2
[10] RColorBrewer_1.1-3 iterators_1.0.14 fastmap_1.2.0
[13] foreach_1.5.2 jsonlite_2.0.0 survival_3.8-3
[16] mgcv_1.9-3 httr_1.4.7 purrr_1.2.0
[19] crosstalk_1.2.2 viridisLite_0.4.2 scales_1.4.0
[22] codetools_0.2-20 abind_1.4-8 cli_3.6.5
[25] rlang_1.3.0 parallelly_1.46.0 future.apply_1.20.1
[28] splines_4.5.2 withr_3.0.2 yaml_2.3.12
[31] otel_0.2.0 tools_4.5.2 parallel_4.5.2
[34] globals_0.18.0 assertthat_0.2.1 vctrs_0.7.3
[37] R6_2.6.1 lifecycle_1.0.5 stringr_1.6.0
[40] htmlwidgets_1.6.4 pkgconfig_2.0.3 pillar_1.11.1
[43] gtable_0.3.6 data.table_1.18.0 glue_1.8.0
[46] tidyselect_1.2.1 xfun_0.55 tibble_3.3.0
[49] knitr_1.51 farver_2.1.2 nlme_3.1-168
[52] htmltools_0.5.9 labeling_0.4.3 rmarkdown_2.30
[55] compiler_4.5.2 S7_0.2.1