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.

library(hal9001)
library(glmnet)
library(furrr)
library(plotly)
library(origami)
library(dplyr)
library(tidyr)
library(ggplot2)

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.

NoteIdentification assumptions

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$Y

Let’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")
Figure 1: Observed outcomes across W2 by treatment status. Horizontal jitter separates the discrete W2 values.

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))
ImportantCross-validation is not cross-fitting

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")
Figure 2: Distribution of fitted treatment probabilities within the observed treatment groups.

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)^2

We 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")
Figure 3: Estimated and true CATE curves over W2 at three fixed values of W1.

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_pred
Figure 4: R-learner estimate of the CATE surface over W1 and W2.

True CATE surface

The true surface combines the plane \(0.5W_1\) with the wave \(\sin(2\pi W_2)\).

plt_true
Figure 5: True CATE surface over W1 and W2.

Reproducibility 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