Conditional density estimation using haldensify

Overview

This coding tutorial uses the haldensify package to estimate a conditional density with the highly adaptive lasso (HAL).

By the end of the tutorial, you will be able to:

  • simulate data from a known conditional distribution;
  • fit a cross-validated HAL conditional density estimator;
  • interpret the empirical-risk path used to select the Lasso penalty;
  • evaluate the fitted density on a structured grid;
  • and compare the estimated and true densities with both line plots and interactive surfaces.

1. Setup

We first load the packages we need for this tutorial. haldensify fits the conditional density; dplyr and tidyr help us inspect and reshape results; ggplot2 creates static diagnostic plots; and plotly creates rotatable three-dimensional surfaces for better interactive visualizations.

# load required packages
library(haldensify)
library(dplyr)
library(tidyr)
library(ggplot2)
library(plotly)

If needed, install the packages before proceeding the tutorial.

install.packages(c("haldensify", "dplyr", "tidyr", "ggplot2", "plotly"))

We set a seed so that the simulated data (and therefore every fitted result in the tutorial) can be reproduced. The helper file sim_data.R lives in the same directory as this document.

set.seed(75681)
source("sim_data.R")

2. Understand the data-generating process

Let \(W=(W_1,W_2)\), where \(W_1,W_2 \overset{\mathrm{iid}}{\sim} N(0,1).\) Conditional on \(W\), the continuous variable \(A\) follows a normal distribution: \(A\mid W=w \sim N\!\left(\mu_A(w),\sigma_A^2(w)\right),\) with \(\mu_A(w)=1+0.3w_1-1.5w_2\) and \(\sigma_A(w)=2+0.3|w_1|.\)

Thus, both the location and spread of \(A\mid W\) change with the covariates. The true conditional density is

\[g_0(a\mid w) =\frac{1}{\sigma_A(w)} \phi\!\left(\frac{a-\mu_A(w)}{\sigma_A(w)}\right),\]

where \(\phi\) is the standard normal density.

Simulate the observed data

We draw \(n=2{,}000\) independent observations. The returned data frame contains one row per observation and columns W1, W2, and A.

n <- 500
data <- sim_data(n)

Here, we check the structure and basic marginal summaries.

dplyr::glimpse(data)
Rows: 500
Columns: 3
$ W1 <dbl> 0.77509790, -1.14246705, 1.63483829, -1.18666555, -1.73437211, 0.98…
$ W2 <dbl> 1.137207387, 1.127008356, 0.120203122, -0.848469904, 1.876121380, -…
$ A  <dbl> -1.3576912, -2.1534798, 3.9102773, 3.3953566, 0.2698453, 4.1968480,…
summary(data)
       W1                 W2                   A          
 Min.   :-2.51634   Min.   :-3.0236447   Min.   :-8.2886  
 1st Qu.:-0.68487   1st Qu.:-0.6892753   1st Qu.:-1.0388  
 Median :-0.07930   Median : 0.0002331   Median : 0.7100  
 Mean   :-0.03675   Mean   : 0.0276948   Mean   : 0.8603  
 3rd Qu.: 0.60739   3rd Qu.: 0.7366232   3rd Qu.: 2.6183  
 Max.   : 3.17026   Max.   : 3.3209255   Max.   : 9.6683  

The next plot shows an exploratory view of the data. The black smooth curve shows how the average of \(A\) varies with \(W_1\), while color reveals the stronger negative shift associated with \(W_2\).

ggplot(data, aes(x = W1, y = A, color = W2)) +
  geom_point(alpha = 0.45, size = 1) +
  geom_smooth(
    aes(group = 1),
    method = "loess",
    se = FALSE,
    color = "black",
    linewidth = 0.9
  ) +
  scale_color_viridis_c() +
  labs(
    x = expression(W[1]),
    y = "A",
    color = expression(W[2]),
    title = "A first look at the simulated data"
  ) +
  theme_minimal(base_size = 12)
Figure 1: Observed values of A against W1, colored by W2. The black curve is a descriptive smooth through the simulated observations.

3. Fit the HAL conditional density estimator

At a high level, haldensify divides the observed support of \(A\) into ordered bins, represents the density through a sequence of conditional hazards, fits those hazards using HAL, and maps the estimated hazards back to a density. Its cross-validation procedure selects tuning parameters by minimizing a density log-likelihood loss.

The fitting call has three key inputs:

  • A = data$A supplies the continuous variable whose density we want;
  • W = data$W supplies the conditioning variables; and
  • lambda_seq supplies candidate Lasso penalties.
haldensify_fit <- haldensify(A = data$A,
                             W = data[, c("W1", "W2")],
                             lambda_seq = exp(seq(-5, -15, length = 50)))

The 50 candidate penalties decrease from \(\exp(-5)\) to \(\exp(-15)\) on a logarithmic scale. Larger values impose more shrinkage; smaller values allow a more flexible HAL fit. The package uses cross-validation to choose among these candidates.

NoteComputational note

This is the most expensive step in the tutorial. It fits a sequence of HAL models inside cross-validation, so runtime depends on the sample size, the number of candidate penalties, and the machine. For quicker experimentation, start with a smaller n or a shorter lambda_seq.

We can inspect the selected penalty and its position in the candidate path.

lambda_summary <- data.frame(
  quantity = c(
    "Largest candidate lambda",
    "Smallest candidate lambda",
    "Cross-validated lambda",
    "Selected path index"
  ),
  value = c(
    formatC(
      max(haldensify_fit$cv_tuning_results$lambda_seq),
      format = "e", digits = 3
    ),
    formatC(
      min(haldensify_fit$cv_tuning_results$lambda_seq),
      format = "e", digits = 3
    ),
    formatC(
      haldensify_fit$cv_tuning_results$lambda_loss_min,
      format = "e", digits = 3
    ),
    sprintf(
      "%d of %d",
      haldensify_fit$cv_tuning_results$lambda_loss_min_idx,
      length(haldensify_fit$cv_tuning_results$lambda_seq)
    )
  )
)

lambda_summary
                   quantity     value
1  Largest candidate lambda 6.738e-03
2 Smallest candidate lambda 3.059e-07
3    Cross-validated lambda 4.746e-04
4       Selected path index  14 of 50

Diagnose the risk path

The package’s plot() method displays empirical risk along the regularization path. Its horizontal axis is \(-\log(\lambda)\), so moving to the right means using a smaller penalty and a more flexible model. The dotted vertical line marks the penalty with minimum cross-validated risk.

# empirical risk vs. lambda
risk_vs_lambda <- plot(haldensify_fit)
risk_vs_lambda
Figure 2: Empirical risk over the HAL regularization path. The dotted line marks the cross-validated choice of lambda.

4. Construct an evaluation grid

The fitted object estimates a function of three arguments, \(g(a\mid w_1,w_2)\). A single three-dimensional surface can display only two input axes plus density height, so we hold \(W_2=0\) and vary \(A\) and \(W_1\):

  • 20 values of \(W_1\) from \(-2\) to \(2\);
  • one fixed value, \(W_2=0\); and
  • 30 values of \(A\) from \(-2\) to \(6\).

The Cartesian product contains \(20\times 1\times 30=600\) evaluation points.

# compute true conditional densities over a grid
W1_vals <- seq(-2, 2, length.out = 20)
W2_fixed <- 0
A_vals <- seq(-2, 6, length.out = 30)
grid <- expand.grid(W1 = W1_vals,
                    W2 = W2_fixed, 
                    A = A_vals)

Inspecting the grid makes its ordering explicit.

dim(grid)
[1] 600   3
head(grid, 8)
          W1 W2  A
1 -2.0000000  0 -2
2 -1.7894737  0 -2
3 -1.5789474  0 -2
4 -1.3684211  0 -2
5 -1.1578947  0 -2
6 -0.9473684  0 -2
7 -0.7368421  0 -2
8 -0.5263158  0 -2
tail(grid, 8)
           W1 W2 A
593 0.5263158  0 6
594 0.7368421  0 6
595 0.9473684  0 6
596 1.1578947  0 6
597 1.3684211  0 6
598 1.5789474  0 6
599 1.7894737  0 6
600 2.0000000  0 6
ImportantThis surface is a slice

The plots below describe \(g(a\mid W_1=w_1,W_2=0)\), not the entire conditional density over every value of \(W_2\). Changing W2_fixed produces a different slice.

5. Evaluate the truth and the fitted density

The simulation helper, with no grid, generates random observations. With a grid, it evaluates the known normal density at each row and returns those density values in its A column. We save them as true_dens.

grid$true_dens <- sim_data(n, grid = grid)$A

Next, we ask the fitted model for one density estimate per grid row. new_A and new_W must describe matching evaluation points and therefore must have the same number of rows.

grid$pred_dens <- predict(haldensify_fit, 
                          new_A = grid$A,
                          new_W = grid[, c("W1", "W2")])

By default, predict.haldensify() uses the cross-validated penalty.

Quantify agreement on this grid

Let’s first first look at some numerical summaries.

density_error <- grid |>
  summarise(
    MAE = mean(abs(pred_dens - true_dens)),
    RMSE = sqrt(mean((pred_dens - true_dens)^2)),
    correlation = cor(pred_dens, true_dens)
  )

density_error
         MAE       RMSE correlation
1 0.01819903 0.02347834   0.9131209

Compare a few one-dimensional slices

We select three representative values of \(W_1\) and overlay the estimated and true density curves.

slice_W1_vals <- W1_vals[c(6, 10, 15)]

grid |>
  filter(W1 %in% slice_W1_vals) |>
  select(A, W1, true_dens, pred_dens) |>
  pivot_longer(
    cols = c(true_dens, pred_dens),
    names_to = "curve",
    values_to = "density"
  ) |>
  mutate(
    curve = recode(
      curve,
      true_dens = "True density",
      pred_dens = "Estimated density"
    ),
    W1_slice = factor(
      sprintf("W1 = %.2f", W1),
      levels = sprintf("W1 = %.2f", slice_W1_vals)
    )
  ) |>
  ggplot(aes(x = A, y = density, color = curve, linetype = curve)) +
  geom_line(linewidth = 0.9) +
  facet_wrap(vars(W1_slice), nrow = 1) +
  scale_color_manual(values = c(
    "Estimated density" = "#238b45",
    "True density" = "#cb181d"
  )) +
  labs(
    x = "A",
    y = "Conditional density",
    color = NULL,
    linetype = NULL
  ) +
  theme_minimal(base_size = 12) +
  theme(legend.position = "bottom")
Figure 3: Estimated and true conditional density curves at three values of W1, with W2 fixed at zero.

Because \(W_2=0\) in these slices, the true mean is \(1+0.3W_1\). As \(W_1\) increases, the center of the density moves to the right. The true standard deviation is \(2+0.3|W_1|\), so the density is also somewhat wider and flatter when \(W_1\) is farther from zero.

6. Build the interactive density surfaces

plotly::add_surface() expects a matrix whose rows correspond to the values on the y axis and whose columns correspond to the values on the x axis. Our plots use \(W_1\) on y and \(A\) on x, so each density matrix must have length(W1_vals) rows and length(A_vals) columns.

This ordering also matches the vector produced from our grid.

pred_dens_mat <- matrix(grid$pred_dens,
                        nrow = length(W1_vals),
                        ncol = length(A_vals),
                        byrow = FALSE)
true_dens_mat <- matrix(grid$true_dens,
                        nrow = length(W1_vals),
                        ncol = length(A_vals),
                        byrow = FALSE)

stopifnot(
  identical(dim(pred_dens_mat), c(length(W1_vals), length(A_vals))),
  identical(dim(true_dens_mat), c(length(W1_vals), length(A_vals)))
)

The estimated density is green and the true density is red. Both use the same axes, camera angle, and vertical range.

plt_pred <- plot_ly(x = A_vals, y = W1_vals, z = pred_dens_mat, 
        colorscale = "Greens") %>%
  add_surface() %>% 
  layout(title = list(text = "Estimated Conditional Density", 
                      font = list(size = 18, color = "black"),
                      yref = "container", 
                      y = 0.98),
         scene = list(xaxis = list(title = "A", titlefont = list(size = 14)),
                      yaxis = list(title = "W1", titlefont = list(size = 14)),
                      zaxis = list(title = "Density", titlefont = list(size = 14), range = c(0, 0.2)),
                      camera = list(eye = list(x = 1.5, y = 1.5, z = 1))),
         legend = list(x = 0.8, y = 1, font = list(size = 12)))

plt_true <- plot_ly(x = A_vals, y = W1_vals, z = true_dens_mat, 
        colorscale = "Reds") %>%
  add_surface() %>%
  layout(title = list(text = "True Conditional Density", 
                      font = list(size = 18, color = "black"),
                      yref = "container", 
                      y = 0.98),
         scene = list(xaxis = list(title = "A", titlefont = list(size = 14)),
                      yaxis = list(title = "W1", titlefont = list(size = 14)),
                      zaxis = list(title = "Density", titlefont = list(size = 14), range = c(0, 0.2)),
                      camera = list(eye = list(x = 1.5, y = 1.5, z = 1))),
         legend = list(x = 0.8, y = 1, font = list(size = 12)))

Estimated conditional density

Drag to rotate the surface, scroll to zoom, and hover to inspect individual values.

plt_pred
Figure 4: HAL estimate of the conditional density of A given W1, evaluated with W2 fixed at zero.

True conditional density

plt_true
Figure 5: True conditional density of A given W1, evaluated with W2 fixed at zero.

7. Adapt the workflow to your own data

For an applied analysis, replace data$A with a numeric continuous variable and replace data$W with a data frame or matrix containing the conditioning variables. The template is:

fit <- haldensify(
  A = observed_A,
  W = observed_W,
  lambda_seq = exp(seq(-0.1, -10, length = 300))
)

density_hat <- predict(
  fit,
  new_A = evaluation_A,
  new_W = evaluation_W
)

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] future_1.68.0    plotly_4.12.1    ggplot2_4.0.3    tidyr_1.3.2     
[5] dplyr_1.2.1      haldensify_0.2.8

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  hal9001_0.4.6       iterators_1.0.14   
[13] fastmap_1.2.0       foreach_1.5.2       jsonlite_2.0.0     
[16] Matrix_1.7-4        glmnet_5.0          survival_3.8-3     
[19] mgcv_1.9-3          httr_1.4.7          origami_1.0.7      
[22] purrr_1.2.0         crosstalk_1.2.2     viridisLite_0.4.2  
[25] scales_1.4.0        codetools_0.2-20    abind_1.4-8        
[28] Rdpack_2.6.4        cli_3.6.5           rlang_1.3.0        
[31] rbibutils_2.4       parallelly_1.46.0   future.apply_1.20.1
[34] splines_4.5.2       withr_3.0.2         yaml_2.3.12        
[37] otel_0.2.0          tools_4.5.2         parallel_4.5.2     
[40] globals_0.18.0      assertthat_0.2.1    vctrs_0.7.3        
[43] R6_2.6.1            matrixStats_1.5.0   lifecycle_1.0.5    
[46] stringr_1.6.0       htmlwidgets_1.6.4   pkgconfig_2.0.3    
[49] pillar_1.11.1       gtable_0.3.6        data.table_1.18.0  
[52] glue_1.8.0          Rcpp_1.1.1          xfun_0.55          
[55] tibble_3.3.0        tidyselect_1.2.1    knitr_1.51         
[58] farver_2.1.2        nlme_3.1-168        htmltools_0.5.9    
[61] labeling_0.4.3      rmarkdown_2.30      compiler_4.5.2     
[64] S7_0.2.1