--- title: "Fitting a proxy to a density you cannot sample" output: rmarkdown::html_vignette vignette: > %\VignetteIndexEntry{Fitting a proxy to a density you cannot sample} %\VignetteEngine{knitr::rmarkdown} %\VignetteEncoding{UTF-8} --- ```{r setup, include = FALSE} knitr::opts_chunk$set( collapse = TRUE, comment = "#>", fig.width = 7, fig.height = 4.5, dpi = 150, out.width = "100%" ) ``` ```{r library} library(proxymix) ``` ```{r engines} has_ggplot2 <- requireNamespace("ggplot2", quietly = TRUE) ``` ```{r stored-results, include = FALSE} ## The comparison table reads stored simulation results. They must come ## from the same major.minor version of proxymix as this build. res <- readRDS("results/quickstart.rds") major_minor <- function(v) paste(unlist(package_version(v))[1:2], collapse = ".") if (major_minor(res$proxymix_version) != major_minor(as.character(packageVersion("proxymix")))) { stop("results/quickstart.rds was built under proxymix ", res$proxymix_version, ", but this is proxymix ", packageVersion("proxymix"), ". Rerun the simulation and ", "data-raw/vignette_results/quickstart.R.", call. = FALSE) } ## Small numbers are written as plain decimals rather than in the ## scientific notation that knitr's inline hook would otherwise use. fixed <- function(v, digits) { format(round(v, digits), nsmall = digits, scientific = FALSE) } ``` ## The problem A statistical distribution, such as the normal, comes with two tools. One is a formula for how likely each value is, which for the normal is the bell curve. The other is a way to generate random values, such as `rnorm()`. With both, you can simulate data, work out means and probabilities, and draw the distribution. In research, often only the first tool is available. A Bayesian analysis, for example, ends with a formula that says how plausible each combination of parameter values is, but gives no direct way to draw from it. You can use the formula to work out how likely any single point is, but it will not give you a random sample, a mean, or the probability of a range of values. proxymix builds a stand-in, or proxy, for such a distribution. The proxy is a mixture of a few normal distributions added together, known as a Gaussian mixture. Normal distributions are easy to work with, so the proxy is too: you can draw from it, average over it, and fix one variable at a value to see how the others behave. The package also reports how close the proxy is to the original, so you know whether to trust it. This vignette works through one example from start to finish. ## Package capabilities - `gmm_target()` describes the distribution you want to approximate, called the target. You supply the number of variables and a function that returns the log of the density. Logs are used because density values can be extremely small. `banana_target()` is a ready-made target in two variables, shaped like a curved banana. - `fit_proxymix()` fits the proxy. It chooses one of three fitting methods from what you supply. - `proposal_mvt()` sets up a wide distribution that is easy to sample, from which the fitting method draws its trial points. - `gmm_fit_quality()` returns a short report on the quality of the fit, called its certificate. - `dgmm()`, `rgmm()`, `gmm_marginalise()` and `gmm_conditionalise()` use the fitted proxy. They give density values, random draws, the distribution of one variable on its own, and the distribution of one variable when another is held at a fixed value. ## Addressing the problem ### Which fitting method applies The package has three ways of fitting the bell curves, described by van der Hoek and Elliott (2024). `fit_proxymix()` picks one according to what you supply. ```{r regime-table, echo = FALSE} regime_tbl <- data.frame( have = c( "Data, and you want one bell curve", "Data, and you want several bell curves", "Only the formula, no data" ), does = c( "Fits one bell curve with the same centre and spread as the data", paste("Moves and reshapes the bell curves, step by step, until together", "they match the data"), paste("Draws trial points, weights each one by the formula, then fits", "the bell curves to the weighted points") ), setting = c("`\"moment\"`", "`\"sample\"`", "`\"kld\"`"), stringsAsFactors = FALSE ) knitr::kable( regime_tbl, col.names = c("What you have", "What the package does", "`regime` setting"), caption = paste( "How `fit_proxymix()` chooses a fitting method. The `regime` argument", "can also name one directly." ) ) ``` If you already have a sample from the target, the first two methods apply, and established packages such as `mclust`, `mixtools` and `flexmix` can also fit the mixture. When you have only the formula, the third method applies. The package was built for this case. ### A target with no sampler The banana target has an exact density formula and no sample, so only the third method applies. ```{r banana} tgt <- banana_target() tgt ``` ### Fit the proxy The third method draws 2,000 trial points from a broad distribution that is easy to sample. Here it is a Student-t distribution, a relative of the normal with heavier tails, made wide enough to cover the banana. Each trial point is then weighted by how much more likely it is under the target than under the broad distribution. The weights do the same job as survey weights that correct an unrepresentative sample. The mixture is refitted to the weighted points in rounds, which stop when the fit no longer improves. The closeness of the proxy to the target is measured by the Kullback-Leibler (KL) divergence. It is zero when the two distributions match and increases as they become more different. The call below asks for a proxy with three components and sets a seed so the result is reproducible. ```{r fit} proposal <- proposal_mvt(n_dim = 2L, mean = c(0, 0), sigma = 4 * diag(2), df = 5) fit <- fit_proxymix(tgt, N = 3L, regime = "kld", proposal = proposal, is_size = 2000L, max_iter = 60L, seed = 1L) fit ``` ### Check the fit before using it Weighted draws have a weakness familiar from survey work. If a handful of respondents carry very large weights, the survey estimate rests on those few people and becomes unstable. The same happens here if a few trial points carry most of the weight. `gmm_fit_quality()` checks for this and for other signs of a poor fit. ```{r certificate} quality <- gmm_fit_quality(fit) ``` ```{r certificate-table, echo = FALSE} cert_tbl <- data.frame( Check = c("fitting method", "rounds settled before the limit", "weights collapsed onto a few draws", "effective sample size", "effective sample size as a share of all draws", "smallest effective sample size of any component", "largest share of the weight held by one draw", "KL divergence on the fitting draws", "half the variance of the log density ratio (a local KL approximation)", "KL divergence on fresh draws", "fresh draws minus fitting draws"), Value = c( quality$regime, as.character(quality$converged), as.character(quality$degenerate), format(round(quality$ess, 1L), nsmall = 1L), format(round(quality$ess_relative, 3L), nsmall = 3L), format(round(quality$min_component_ess, 1L), nsmall = 1L), format(signif(quality$max_weight, 3L)), format(signif(quality$kld_final, 3L)), format(signif(quality$kld_approx, 3L)), format(signif(quality$heldout_kld, 3L)), format(signif(quality$validation_gap, 3L)) ), stringsAsFactors = FALSE ) knitr::kable( cert_tbl, caption = "The fit certificate returned by `gmm_fit_quality()`." ) ``` The effective sample size is the number of equally weighted draws that the weighted sample is worth. The fit was tuned to the draws it was fitted on, so the KL computed on them is too low. The KL on a fresh set of draws is the one to report. The package flags a fit when this KL exceeds 0.3, when the fit did not converge, or when it is degenerate. Half the variance of the log density ratio on the fitting draws approximates the KL when the proxy is already close to the target. It is only a rough check. ### Compare the proxy with the target With two variables, the quickest check is a plot. ```{r overlay-grid} grid_x <- seq(-3, 3, length.out = 120L) grid_g <- expand.grid(x1 = grid_x, x2 = grid_x) grid_mat <- as.matrix(grid_g) grid_g$target <- exp(tgt@log_density(grid_mat)) grid_g$proxy <- dgmm(grid_mat, fit) ``` ```{r overlay, eval = has_ggplot2, echo = has_ggplot2, fig.cap = sprintf("The banana target (filled contours) with the three-component proxy overlaid as dashed contours. The dashed contours follow the curve of the banana, which a single bell curve could not do. The KL divergence on fresh draws is %s.", format(signif(fit@diagnostics$validation_kld, 2L))), fig.alt = "Filled contour map of the curved banana density with dashed contours of the three-component Gaussian-mixture proxy following the same curve."} ggplot2::ggplot(grid_g, ggplot2::aes(x1, x2)) + ggplot2::geom_contour_filled(ggplot2::aes(z = target), bins = 10L, alpha = 0.85) + ggplot2::geom_contour(ggplot2::aes(z = proxy), colour = "white", linetype = "dashed", linewidth = 0.45, bins = 8L) + ggplot2::scale_fill_viridis_d(option = "mako", guide = "none") + ggplot2::coord_equal() + ggplot2::labs( title = "Target (filled) and fitted proxy (dashed)", x = expression(x[1]), y = expression(x[2]) ) + ggplot2::theme_minimal(base_size = 11) ``` ```{r overlay-skip, eval = !has_ggplot2, echo = FALSE, results = "asis"} cat("ggplot2 is not installed on this build, so the target-and-proxy", "overlay figure is skipped.\n") ``` ### Use the proxy Questions that were hard to answer for the target have exact answers for the proxy, because it is built from normal distributions. `gmm_marginalise(keep = 1L)` gives the distribution of the first variable on its own. `gmm_conditionalise(given = c(NA, 0.5))` gives the distribution of the first variable when the second equals 0.5, with `NA` marking the variable left free. Neither call goes back to the target formula. ```{r operations} gmm_marginalise(fit, keep = 1L) gmm_conditionalise(fit, given = c(NA, 0.5)) ``` Drawing from the proxy is fast. ```{r sample} draws <- rgmm(500L, fit) dim(draws) ``` ### Comparison with the Laplace approximation, Stan and BayesianTools ```{r compare-facts, include = FALSE} sim_value <- function(method, what) { res$sim_tab[[what]][res$sim_tab$method == method] } kl_value <- function(method) fixed(sim_value(method, "kl_mean"), 4) tail_error <- function(method) fixed(sim_value(method, "tail_rmse"), 4) secs <- function(method) fixed(sim_value(method, "secs"), 2) tail_others <- vapply(c("NUTS", "NUTS + mclust", "DEzs"), sim_value, numeric(1L), what = "tail_rmse") tail_range <- paste(fixed(min(tail_others), 4), "to", fixed(max(tail_others), 4)) # the three-component fit from above, scored on the simulation's grid quad_grid <- as.matrix(expand.grid(x1 = seq(-5 + res$h / 2, 5, by = res$h), x2 = seq(-5 + res$h / 2, 12, by = res$h))) quad_log_f <- tgt@log_density(quad_grid) kl_fit_quad <- sum(exp(quad_log_f) * (quad_log_f - dgmm(quad_grid, fit, log = TRUE))) * res$h^2 ``` In a simulation, proxymix was compared with three established methods for a distribution that is known only by its formula. The Laplace approximation (Tierney and Kadane, 1986) is a single normal distribution centred on the highest point of the target, with a spread set by how sharply the target falls away from that point. Stan (Carpenter et al., 2017), run from R through `cmdstanr`, draws a sample from the formula with the no-U-turn sampler, or NUTS (Hoffman and Gelman, 2014). NUTS is a Markov chain Monte Carlo method: it produces random values by a long chain of small, linked steps. A mixture was then fitted to the NUTS draws with `mclust` (Scrucca et al., 2016). DEzs (ter Braak and Vrugt, 2008), from `BayesianTools` (Hartig et al., 2026), is another Markov chain Monte Carlo method. Each method was run `r res$n_rep` times on the banana target without further tuning. proxymix chose its number of components automatically with `select_N()`, and `mclust` chose its own from 1 to 6. Only this one target in two variables was used, and the results may not carry over to targets with more variables. The first measure is the KL divergence of each fitted density from the target, computed by summing over a fine grid of points. It cannot be computed for the two samplers, which return draws but no density. The second measure is the error in the estimated probability that $x_1$ is greater than 2, which is `r fixed(res$tail_ref, 4)` for the target. Smaller is better for both. ```{r compare-table, echo = FALSE} methods <- c("proxymix", "NUTS + mclust", "Laplace", "NUTS", "DEzs") cmp_tbl <- data.frame( method = c("proxymix", "NUTS draws, then mclust", "Laplace approximation", "NUTS draws", "DEzs draws"), kl = vapply(methods, function(s1) { if (is.na(sim_value(s1, "kl_mean"))) return("--") paste0(kl_value(s1), " (", fixed(sim_value(s1, "kl_sd"), 4), ")") }, character(1L)), tail = vapply(methods, function(s1) { fixed(sim_value(s1, "tail_rmse"), 5) }, character(1L)), secs = vapply(methods, function(s1) { if (sim_value(s1, "secs") < 0.01) return("< 0.01") secs(s1) }, character(1L)), stringsAsFactors = FALSE ) knitr::kable( cmp_tbl, row.names = FALSE, align = c("l", "r", "r", "r"), col.names = c("Method", "KL divergence, mean (sd)", "Error in $P(x_1 > 2)$", "Seconds per run"), caption = paste0( "Results over ", res$n_rep, " runs of each method on the banana ", "target. The KL divergence is averaged over the runs, with its ", "standard deviation in brackets. The error in $P(x_1 > 2)$ is the root ", "mean squared error over the runs. Seconds per run is the median; the ", "proxymix time includes choosing the number of components, and the ", "NUTS times leave out the one-off compilation of the Stan program." ) ) ``` proxymix gave the smallest KL divergence, `r kl_value("proxymix")` against `r kl_value("NUTS + mclust")` for the mixture fitted to the NUTS draws. The Laplace approximation was far behind on this measure, at `r kl_value("Laplace")`, because one bell curve cannot follow the curve of the banana. On the tail probability, however, the Laplace approximation was the most accurate, with an error of `r fixed(sim_value("Laplace", "tail_rmse"), 5)`. On its own, $x_1$ has a standard normal distribution, and on this target the Laplace approximation reproduces it almost exactly. proxymix came second on the tail, with an error of `r tail_error("proxymix")` against `r tail_range` for the two samplers and the mixture fitted to the NUTS draws. The held-out KL of `r fixed(fit@diagnostics$validation_kld, 4)` reported for the three-component fit above is itself estimated from random draws, with a standard error of about `r fixed(fit@diagnostics$validation_mc_se, 4)`. On the grid used for the table, that fit has a KL divergence of `r fixed(kl_fit_quad, 4)`. That value lies `r fixed((kl_fit_quad - sim_value("proxymix", "kl_mean")) / sim_value("proxymix", "kl_sd"), 1)` standard deviations above the proxymix mean in the table, using the standard deviation across runs shown there in brackets. The Laplace approximation was also the fastest, at `r if (sim_value("Laplace", "secs") < 0.01) "under 0.01" else secs("Laplace")` seconds per run. DEzs took `r secs("DEzs")` seconds per run and was slightly faster than proxymix at `r secs("proxymix")`, while NUTS took `r secs("NUTS")` and NUTS followed by `mclust` took `r secs("NUTS + mclust")` (medians on one computer, with other programs running on it at the same time). The code below runs each method once and scores it against the grid. It needs `mclust` and `BayesianTools` from CRAN, and `cmdstanr` from with a CmdStan installation. It is not run when this vignette is built. ```{r compare-code, eval = FALSE} library(proxymix) library(cmdstanr) library(mclust) library(BayesianTools) tgt <- banana_target() # midpoint-rule quadrature; the target's mass outside the box is below 1e-5 h <- 0.05 grid <- as.matrix(expand.grid(x1 = seq(-5 + h / 2, 5, by = h), x2 = seq(-5 + h / 2, 12, by = h))) log_f <- tgt@log_density(grid) f_grid <- exp(log_f) tail_ref <- sum(f_grid[grid[, 1L] > 2]) * h^2 kl_of <- function(log_g) sum(f_grid * (log_f - log_g)) * h^2 # mass above x1 = 2 under a Gaussian mixture, from its x1 marginal tail_of <- function(w, mean1, sd1) { sum(w * pnorm(2, mean1, sd1, lower.tail = FALSE)) } # the same target as a Stan program stan_file <- file.path(tempdir(), "banana.stan") writeLines(c( "parameters {", " vector[2] x;", "}", "model {", " target += -0.5 * (square(x[1])", " + square(x[2] - 0.5 * (square(x[1]) - 1)));", "}"), stan_file) banana_stan <- cmdstan_model(stan_file) box <- createBayesianSetup(likelihood = function(x) tgt@log_density(x), lower = c(-6, -6), upper = c(6, 14)) set.seed(1L) fit <- select_N(tgt, seed = 1L)$best_fit fit_1 <- gmm_marginalise(fit, keep = 1L) la <- optim(c(0, 0), function(x) -tgt@log_density(x), method = "BFGS", hessian = TRUE) la_mean <- la$par la_cov <- solve(la$hessian) nuts <- banana_stan$sample(seed = 1L, refresh = 0L, show_messages = FALSE) draws <- nuts$draws("x", format = "matrix") mc <- Mclust(draws, G = 1:6, modelNames = "VVV", verbose = FALSE) mc_par <- mc$parameters de <- runMCMC(box, sampler = "DEzs", settings = list(message = FALSE)) de_draws <- getSample(de, start = 1000L) # KL divergence from the target, for the methods that return a density c(proxymix = kl_of(dgmm(grid, fit, log = TRUE)), Laplace = kl_of(dmvnorm(grid, la_mean, la_cov, log = TRUE)), "NUTS + mclust" = kl_of(dens(grid, mc$modelName, parameters = mc_par, logarithm = TRUE))) # error in the estimated probability that x1 > 2 c(proxymix = tail_of(gmm_weights(fit_1), vapply(gmm_means(fit_1), `[[`, numeric(1L), 1L), sqrt(vapply(gmm_covariances(fit_1), `[[`, numeric(1L), 1L))), Laplace = pnorm(2, la_mean[1L], sqrt(la_cov[1L, 1L]), lower.tail = FALSE), NUTS = mean(draws[, 1L] > 2), "NUTS + mclust" = tail_of(mc_par$pro, mc_par$mean[1L, ], sqrt(mc_par$variance$sigma[1L, 1L, ])), DEzs = mean(de_draws[, 1L] > 2)) - tail_ref ``` The [extended version of this article](https://max578.github.io/proxymix/articles/extended/quickstart.html) gives the full simulation, including how many settings each method needs the user to choose. ## Interpretation The proxy is a mixture of `r gmm_n_components(fit)` normal distributions in `r gmm_dim(fit)` variables, fitted in `r length(kld_trace(fit))` rounds. It can be sampled and summarised without calling the target formula again. The figure shows why a mixture is needed: one bell curve cannot follow a curved shape, but three placed along the curve can. The certificate is consistent with a good fit. The rounds settled before the limit of 60, and the weights did not collapse. The 2,000 weighted draws were worth `r round(fit@diagnostics$ess, 0L)` equally weighted draws, `r round(100 * fit@diagnostics$ess_relative, 0L)` per cent of the total. The heaviest single draw held `r format(signif(100 * fit@diagnostics$max_weight, 1L))` per cent of the total weight, so no single draw dominated the fit. Each component was estimated from an effective sample of at least `r round(quality$min_component_ess, 0L)` draws. The KL divergence on fresh draws is `r format(signif(fit@diagnostics$validation_kld, 2L))`, with a standard error of `r format(signif(fit@diagnostics$validation_mc_se, 2L))`. For a sense of scale, this value means that the probabilities the proxy and the target give to any region differ by at most $\sqrt{\mathrm{KL}/2}$, about `r round(100 * sqrt(fit@diagnostics$validation_kld / 2), 0L)` percentage points (Pinsker's inequality). The KL divergence computed on the grid, `r fixed(kl_fit_quad, 4)`, lies `r if (kl_fit_quad > fit@diagnostics$validation_kld) "above" else "below"` the fresh-draw estimate by `r fixed(abs(kl_fit_quad - fit@diagnostics$validation_kld) / fit@diagnostics$validation_mc_se, 1)` times the standard error of that estimate. The KL on the fitting draws is lower because the fit was tuned to those draws. The difference between the two, `r format(signif(quality$validation_gap, 2L))`, is about `r round(quality$validation_gap / fit@diagnostics$mc_se_kld, 0L)` times the standard error of the KL on the fitting draws, which is `r format(signif(fit@diagnostics$mc_se_kld, 2L))`. ## Limitations The number of components is set to three by hand here. `select_N()` chooses it automatically, and `bic_aic()` reports the BIC and AIC for comparing counts. Too few components show up as a KL divergence that more trial draws do not reduce. The choice of broad distribution matters. The Student-t used here is wide enough to cover the banana. If the broad distribution misses part of the target, the proxy misses that part too, even though the printed mixture may look reasonable. This example has two variables. Weighted trial draws lose efficiency quickly as the number of variables grows, so a proxy of the same quality in five or ten variables needs many more draws. The effective sample size in the certificate shows when this happens. ## Further reading *Choosing between the three fitting regimes* runs all three fitting methods on a target whose true shape is known, so the cost of the wrong choice is visible. *How well a mixture proxies four awkward shapes* applies the third method to a curved ridge, a ring, two well-separated clusters and a distribution with hard edges, and shows the package refusing a fit whose weights have collapsed. *The closed-form operator calculus on a mixture* goes further with the exact operations shown above. *Compressing a Bayesian posterior you can evaluate but not sample* applies this workflow to the result of a Bayesian analysis. ## References Carpenter, B., Gelman, A., Hoffman, M. D., Lee, D., Goodrich, B., Betancourt, M., Brubaker, M., Guo, J., Li, P. and Riddell, A. (2017). *Stan: A probabilistic programming language.* Journal of Statistical Software 76(1), 1--32. . Hartig, F., Minunno, F. and Paul, S. (2026). *BayesianTools: General-purpose MCMC and SMC samplers and tools for Bayesian statistics.* R package version `r res$versions[["BayesianTools"]]`. . Hoek, J. van der and Elliott, R. J. (2024). *Mixtures of multivariate Gaussians.* Stochastic Analysis and Applications. . Hoffman, M. D. and Gelman, A. (2014). *The No-U-Turn sampler: Adaptively setting path lengths in Hamiltonian Monte Carlo.* Journal of Machine Learning Research 15(47), 1593--1623. . Scrucca, L., Fop, M., Murphy, T. B. and Raftery, A. E. (2016). *mclust 5: Clustering, classification and density estimation using Gaussian finite mixture models.* The R Journal 8(1), 289--317. . ter Braak, C. J. F. and Vrugt, J. A. (2008). *Differential evolution Markov chain with snooker updater and fewer chains.* Statistics and Computing 18(4), 435--446. . Tierney, L. and Kadane, J. B. (1986). *Accurate approximations for posterior moments and marginal densities.* Journal of the American Statistical Association 81(393), 82--86. . ## Reproduce Every fit is seeded (`seed = 1L`), so re-running this vignette reproduces the same numbers. ```{r session-info, collapse = FALSE, class.output = "session-info"} sessionInfo() ```