# load required packages
library(haldensify)
library(dplyr)
library(tidyr)
library(ggplot2)
library(plotly)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.
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)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$Asupplies the continuous variable whose density we want;W = data$Wsupplies the conditioning variables; andlambda_seqsupplies 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.
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_lambda4. 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
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)$ANext, 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")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_predTrue conditional density
plt_true7. 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