---
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}
---

<!-- Render time: ~2 s under rmarkdown::render() with ggplot2 installed;
     macOS arm64 (Apple silicon), R 4.5.2, one core. The comparison with
     the Laplace approximation, Stan and BayesianTools is read from
     results/quickstart.rds, built by data-raw/vignette_results/quickstart.R. -->

```{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
<https://stan-dev.r-universe.dev> 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. <https://doi.org/10.18637/jss.v076.i01>.

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"]]`.
<https://doi.org/10.32614/CRAN.package.BayesianTools>.

Hoek, J. van der and Elliott, R. J. (2024). *Mixtures of multivariate
Gaussians.* Stochastic Analysis and Applications.
<https://doi.org/10.1080/07362994.2024.2372605>.

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.
<https://jmlr.org/papers/v15/hoffman14a.html>.

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.
<https://doi.org/10.32614/RJ-2016-021>.

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. <https://doi.org/10.1007/s11222-008-9104-9>.

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.
<https://doi.org/10.1080/01621459.1986.10478240>.

## 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()
```
