Adaptive Rejection Sampling
#| '!! shinylive warning !!': |
#| shinylive does not work in self-contained HTML documents.
#| Please set `embed-resources: false` in your metadata.
#| standalone: true
#| viewerHeight: 860
library(shiny)
h <- function(x) -x^2 / 2 # log unnormalized N(0,1) density
hprime <- function(x) -x
build_hull <- function(xs, lb, ub) {
xs <- sort(xs)
hx <- h(xs); hpx <- hprime(xs)
k <- length(xs)
z <- numeric(k + 1); z[1] <- lb; z[k + 1] <- ub
if (k > 1) for (i in 1:(k - 1)) {
z[i + 1] <- (hx[i + 1] - hx[i] - xs[i + 1] * hpx[i + 1] + xs[i] * hpx[i]) /
(hpx[i] - hpx[i + 1])
}
list(x = xs, hx = hx, hp = hpx, z = z)
}
### log-area under the exponential of a single tangent-line piece, between its
### breakpoints -- computed entirely in log space so it stays finite even far
### out in the tail, where exp(hull value) itself would underflow to exactly 0
log_piece_area <- function(hull, j) {
a <- hull$hx[j]; b <- hull$hp[j]; xj <- hull$x[j]
lo <- hull$z[j]; hi <- hull$z[j + 1]
val_lo <- a + b * (lo - xj); val_hi <- a + b * (hi - xj)
if (abs(b) < 1e-10) return(val_lo + log(hi - lo))
m <- max(val_lo, val_hi); d <- min(val_lo, val_hi)
m + log1p(-exp(d - m)) - log(abs(b))
}
### draw one point from the piecewise-exponential envelope: pick a piece
### proportional to its area, then invert its (exponential) CDF within it --
### both steps done on the log scale for the same underflow-safety reason
sample_hull <- function(hull) {
k <- length(hull$x)
log_areas <- vapply(seq_len(k), function(j) log_piece_area(hull, j), numeric(1))
log_areas[!is.finite(log_areas)] <- -Inf
areas <- exp(log_areas - max(log_areas)) # relative areas, safely in (0, 1]
j <- sample.int(k, 1, prob = areas)
a <- hull$hx[j]; b <- hull$hp[j]; xj <- hull$x[j]
lo <- hull$z[j]; hi <- hull$z[j + 1]
U <- runif(1)
if (abs(b) < 1e-10) {
x <- lo + U * (hi - lo)
logg <- a
} else {
val_lo <- a + b * (lo - xj); val_hi <- a + b * (hi - xj)
m <- max(val_lo, val_hi)
val_x <- m + log((1 - U) * exp(val_lo - m) + U * exp(val_hi - m))
x <- xj + (val_x - a) / b
logg <- val_x
}
list(x = x, logg = logg)
}
### piecewise-linear hull value (log scale) at a vector of x's
hull_val <- function(hull, xgrid) {
j <- findInterval(xgrid, hull$z, rightmost.closed = TRUE)
j <- pmax(1, pmin(j, length(hull$x)))
hull$hx[j] + hull$hp[j] * (xgrid - hull$x[j])
}
init_hull <- function(lb, ub) {
x_init <- sort(unique(pmin(pmax(c(lb + 0.25 * (ub - lb), (lb + ub) / 2,
ub - 0.25 * (ub - lb)), lb + 1e-6), ub - 1e-6)))
build_hull(x_init, lb, ub)
}
### smart plotting window: view = [mode - 2, mode + 2], clipped to [lb, ub],
### where "mode" is the peak of the density *within* the truncation window --
### 0 if the window straddles it, otherwise whichever edge is nearest 0 (the
### window is then entirely on one side of the peak, so the density is
### monotonic and highest right at that edge). A wide window far out in the
### tail (e.g. [100, 200]) would otherwise dilute the whole plot into a
### mostly-flat, uninformative sliver near that one edge.
smart_xlim <- function(lb, ub) {
m <- if (lb >= 0) lb else if (ub <= 0) ub else 0
c(max(lb, m - 2), min(ub, m + 2))
}
### log(pnorm(ub) - pnorm(lb)), numerically stable even when the naive
### difference underflows to 0 for windows far out in either tail
log_pnorm_diff <- function(lb, ub) {
if (ub <= 0) {
log_pu <- pnorm(ub, log.p = TRUE); log_pl <- pnorm(lb, log.p = TRUE)
log_pu + log1p(-exp(log_pl - log_pu))
} else if (lb >= 0) {
log_su <- pnorm(ub, lower.tail = FALSE, log.p = TRUE)
log_sl <- pnorm(lb, lower.tail = FALSE, log.p = TRUE)
log_sl + log1p(-exp(log_su - log_sl))
} else {
log(pnorm(ub) - pnorm(lb))
}
}
tnorm_density <- function(x, lb, ub) exp(dnorm(x, log = TRUE) - log_pnorm_diff(lb, ub))
ui <- fluidPage(
titlePanel("Shinylive App for Adaptive Rejection Sampling"),
sidebarLayout(
sidebarPanel(
width = 3,
numericInput("n", "sample size (target)", value = 200, min = 1, step = 10),
numericInput("l", "lower limit l", value = -2, step = 0.1),
numericInput("u", "upper limit u", value = 2, step = 0.1),
fluidRow(
column(6, actionButton("play_pause", "▶ Play", width = "100%")),
column(6, actionButton("next_step", "Next ↷", width = "100%"))
),
br(),
actionButton("restart", "↺ Restart", width = "100%"),
helpText("Each step draws one candidate from the current hull. ",
"Accepted points (green) join the histogram; rejected ",
"points (black x) become new tangent lines, tightening ",
"the hull. Changing the limits restarts the run."),
textOutput("status")
),
mainPanel(
width = 9,
plotOutput("panels", height = "620px")
)
)
)
server <- function(input, output, session) {
rv <- reactiveValues(hull = NULL, samples = numeric(0), tries = 0, target_n = 200,
lb = -2, ub = 2, done = FALSE, playing = FALSE,
last_x = NULL, last_accept = NA)
reset_run <- function() {
lb <- if (!is.null(input$l) && !is.na(input$l)) input$l else -2
ub <- if (!is.null(input$u) && !is.na(input$u)) input$u else 2
if (ub - lb < 0.05) ub <- lb + 0.05 # guard against a degenerate window
rv$hull <- init_hull(lb, ub)
rv$samples <- numeric(0)
rv$tries <- 0
rv$lb <- lb; rv$ub <- ub
rv$target_n <- if (!is.null(input$n) && !is.na(input$n) && input$n >= 1) input$n else 200
rv$done <- FALSE
rv$playing <- FALSE
rv$last_x <- NULL; rv$last_accept <- NA
updateActionButton(session, "play_pause", label = "▶ Play")
}
isolate(reset_run())
do_step <- function() {
if (rv$done) return(invisible(NULL))
draw <- sample_hull(rv$hull)
x <- draw$x; logg <- draw$logg; logf <- h(x)
rv$tries <- rv$tries + 1
accept <- log(runif(1)) < logf - logg
rv$last_x <- x; rv$last_accept <- accept
if (accept) {
rv$samples <- c(rv$samples, x)
} else {
rv$hull <- build_hull(sort(c(rv$hull$x, x)), rv$lb, rv$ub)
}
if (length(rv$samples) >= rv$target_n) {
rv$done <- TRUE
rv$playing <- FALSE
updateActionButton(session, "play_pause", label = "▶ Play")
}
}
observeEvent(input$l, { reset_run() }, ignoreInit = TRUE)
observeEvent(input$u, { reset_run() }, ignoreInit = TRUE)
observeEvent(input$restart, { reset_run() })
observeEvent(input$n, {
req(input$n)
rv$target_n <- input$n
if (length(rv$samples) < rv$target_n) rv$done <- FALSE
}, ignoreInit = TRUE)
observeEvent(input$play_pause, {
if (rv$done) return()
rv$playing <- !rv$playing
updateActionButton(session, "play_pause",
label = if (rv$playing) "⏸ Pause" else "▶ Play")
})
observeEvent(input$next_step, {
rv$playing <- FALSE
updateActionButton(session, "play_pause", label = "▶ Play")
do_step()
})
### auto-advance one step every 120ms while playing
observe({
if (isTRUE(rv$playing) && !rv$done) {
invalidateLater(120, session)
isolate(do_step())
}
})
output$status <- renderText({
sprintf("%d / %d samples accepted | %d proposals so far | %d tangent points",
length(rv$samples), rv$target_n, rv$tries, length(rv$hull$x))
})
output$panels <- renderPlot({
hull <- rv$hull; lb <- rv$lb; ub <- rv$ub
dl <- smart_xlim(lb, ub)
xg <- seq(dl[1], dl[2], length.out = 400)
hull_log <- hull_val(hull, xg); target_log <- h(xg)
par(mfrow = c(2, 1), mar = c(4, 4, 3, 1))
plot(xg, hull_log, type = "l", lwd = 2, col = "steelblue",
ylim = range(c(hull_log, target_log)), xlab = "x",
ylab = "log density (unnormalized)",
main = sprintf("log target vs. hull (%d tangent points)", length(hull$x)))
lines(xg, target_log, lwd = 2, col = "firebrick")
points(hull$x, hull$hx, pch = 19, col = "gray20", cex = 0.9)
abline(v = hull$z[-c(1, length(hull$z))], col = adjustcolor("gray40", 0.4), lty = 3)
abline(v = c(lb, ub), col = "darkorange", lty = 2)
if (!is.null(rv$last_x)) {
points(rv$last_x, h(rv$last_x), pch = if (rv$last_accept) 19 else 4,
col = if (rv$last_accept) "forestgreen" else "black", cex = 1.6, lwd = 2)
}
legend("bottomleft", bty = "n", cex = 0.85,
lwd = c(2, 2, 1, NA, NA), lty = c(1, 1, 2, NA, NA), pch = c(NA, NA, NA, 19, 4),
col = c("steelblue", "firebrick", "darkorange", "forestgreen", "black"),
legend = c("log hull", "log target", "truncation limits", "last: accepted", "last: rejected"))
if (length(rv$samples) == 0) {
plot(0, 0, type = "n", xlim = dl, ylim = c(0, 1), xlab = "x", ylab = "",
main = "no samples accepted yet")
} else {
### scale the bin count so ~30 bins fall within the *visible* view,
### not the full [lb, ub] window -- otherwise a wide window (e.g.
### [100, 200] zoomed to [100, 102]) puts everything in one giant bin
n_bins <- min(5000, max(30, ceiling(30 * (ub - lb) / (dl[2] - dl[1]))))
breaks <- seq(lb, ub, length.out = n_bins + 1)
hist(rv$samples, breaks = breaks, freq = FALSE, col = "steelblue", border = "white",
xlim = dl, xlab = "x",
main = sprintf("%d / %d samples%s", length(rv$samples), rv$target_n,
if (rv$done) " (done)" else ""))
curve(tnorm_density(x, lb, ub), add = TRUE, lwd = 3, col = "firebrick", n = 300)
legend("topright", bty = "n", lwd = 3, col = "firebrick", legend = "true truncated-normal density")
}
})
}
shinyApp(ui, server)
About the app
Shows how adaptive rejection sampling improves its envelope automatically, so that it becomes more efficient as sampling proceeds. The app above reimplements the same idea directly — start with a few tangent lines to \(\log f\), sample from the piecewise-exponential envelope they define, and whenever a candidate is rejected, add it as a new tangent point and rebuild the hull — one candidate at a time, so every step of the mechanism is visible. Next advances a single candidate; Play advances automatically.
The ars package above is a black box: it returns samples but not the sequence of hulls it built along the way. The top panel shows \(\log f\) against the current hull (in log scale the hull is exactly piecewise-linear, since it is built from tangent lines), marking the most recent candidate as accepted (green dot) or rejected (black x, becoming a new tangent point). The bottom panel is a histogram of the accepted draws so far against the true truncated normal density. Controls: the target sample size, and the truncation limits \([l, u]\) (changing the limits restarts the run).
library(shiny)
h <- function(x) -x^2 / 2 # log unnormalized N(0,1) density
hprime <- function(x) -x
build_hull <- function(xs, lb, ub) {
xs <- sort(xs)
hx <- h(xs); hpx <- hprime(xs)
k <- length(xs)
z <- numeric(k + 1); z[1] <- lb; z[k + 1] <- ub
if (k > 1) for (i in 1:(k - 1)) {
z[i + 1] <- (hx[i + 1] - hx[i] - xs[i + 1] * hpx[i + 1] + xs[i] * hpx[i]) /
(hpx[i] - hpx[i + 1])
}
list(x = xs, hx = hx, hp = hpx, z = z)
}
### log-area under the exponential of a single tangent-line piece, between its
### breakpoints -- computed entirely in log space so it stays finite even far
### out in the tail, where exp(hull value) itself would underflow to exactly 0
log_piece_area <- function(hull, j) {
a <- hull$hx[j]; b <- hull$hp[j]; xj <- hull$x[j]
lo <- hull$z[j]; hi <- hull$z[j + 1]
val_lo <- a + b * (lo - xj); val_hi <- a + b * (hi - xj)
if (abs(b) < 1e-10) return(val_lo + log(hi - lo))
m <- max(val_lo, val_hi); d <- min(val_lo, val_hi)
m + log1p(-exp(d - m)) - log(abs(b))
}
### draw one point from the piecewise-exponential envelope: pick a piece
### proportional to its area, then invert its (exponential) CDF within it --
### both steps done on the log scale for the same underflow-safety reason
sample_hull <- function(hull) {
k <- length(hull$x)
log_areas <- vapply(seq_len(k), function(j) log_piece_area(hull, j), numeric(1))
log_areas[!is.finite(log_areas)] <- -Inf
areas <- exp(log_areas - max(log_areas)) # relative areas, safely in (0, 1]
j <- sample.int(k, 1, prob = areas)
a <- hull$hx[j]; b <- hull$hp[j]; xj <- hull$x[j]
lo <- hull$z[j]; hi <- hull$z[j + 1]
U <- runif(1)
if (abs(b) < 1e-10) {
x <- lo + U * (hi - lo)
logg <- a
} else {
val_lo <- a + b * (lo - xj); val_hi <- a + b * (hi - xj)
m <- max(val_lo, val_hi)
val_x <- m + log((1 - U) * exp(val_lo - m) + U * exp(val_hi - m))
x <- xj + (val_x - a) / b
logg <- val_x
}
list(x = x, logg = logg)
}
### piecewise-linear hull value (log scale) at a vector of x's
hull_val <- function(hull, xgrid) {
j <- findInterval(xgrid, hull$z, rightmost.closed = TRUE)
j <- pmax(1, pmin(j, length(hull$x)))
hull$hx[j] + hull$hp[j] * (xgrid - hull$x[j])
}
init_hull <- function(lb, ub) {
x_init <- sort(unique(pmin(pmax(c(lb + 0.25 * (ub - lb), (lb + ub) / 2,
ub - 0.25 * (ub - lb)), lb + 1e-6), ub - 1e-6)))
build_hull(x_init, lb, ub)
}
### smart plotting window: view = [mode - 2, mode + 2], clipped to [lb, ub],
### where "mode" is the peak of the density *within* the truncation window --
### 0 if the window straddles it, otherwise whichever edge is nearest 0 (the
### window is then entirely on one side of the peak, so the density is
### monotonic and highest right at that edge). A wide window far out in the
### tail (e.g. [100, 200]) would otherwise dilute the whole plot into a
### mostly-flat, uninformative sliver near that one edge.
smart_xlim <- function(lb, ub) {
m <- if (lb >= 0) lb else if (ub <= 0) ub else 0
c(max(lb, m - 2), min(ub, m + 2))
}
### log(pnorm(ub) - pnorm(lb)), numerically stable even when the naive
### difference underflows to 0 for windows far out in either tail
log_pnorm_diff <- function(lb, ub) {
if (ub <= 0) {
log_pu <- pnorm(ub, log.p = TRUE); log_pl <- pnorm(lb, log.p = TRUE)
log_pu + log1p(-exp(log_pl - log_pu))
} else if (lb >= 0) {
log_su <- pnorm(ub, lower.tail = FALSE, log.p = TRUE)
log_sl <- pnorm(lb, lower.tail = FALSE, log.p = TRUE)
log_sl + log1p(-exp(log_su - log_sl))
} else {
log(pnorm(ub) - pnorm(lb))
}
}
tnorm_density <- function(x, lb, ub) exp(dnorm(x, log = TRUE) - log_pnorm_diff(lb, ub))
ui <- fluidPage(
titlePanel("Shinylive App for Adaptive Rejection Sampling"),
sidebarLayout(
sidebarPanel(
width = 3,
numericInput("n", "sample size (target)", value = 200, min = 1, step = 10),
numericInput("l", "lower limit l", value = -2, step = 0.1),
numericInput("u", "upper limit u", value = 2, step = 0.1),
fluidRow(
column(6, actionButton("play_pause", "▶ Play", width = "100%")),
column(6, actionButton("next_step", "Next ↷", width = "100%"))
),
br(),
actionButton("restart", "↺ Restart", width = "100%"),
helpText("Each step draws one candidate from the current hull. ",
"Accepted points (green) join the histogram; rejected ",
"points (black x) become new tangent lines, tightening ",
"the hull. Changing the limits restarts the run."),
textOutput("status")
),
mainPanel(
width = 9,
plotOutput("panels", height = "620px")
)
)
)
server <- function(input, output, session) {
rv <- reactiveValues(hull = NULL, samples = numeric(0), tries = 0, target_n = 200,
lb = -2, ub = 2, done = FALSE, playing = FALSE,
last_x = NULL, last_accept = NA)
reset_run <- function() {
lb <- if (!is.null(input$l) && !is.na(input$l)) input$l else -2
ub <- if (!is.null(input$u) && !is.na(input$u)) input$u else 2
if (ub - lb < 0.05) ub <- lb + 0.05 # guard against a degenerate window
rv$hull <- init_hull(lb, ub)
rv$samples <- numeric(0)
rv$tries <- 0
rv$lb <- lb; rv$ub <- ub
rv$target_n <- if (!is.null(input$n) && !is.na(input$n) && input$n >= 1) input$n else 200
rv$done <- FALSE
rv$playing <- FALSE
rv$last_x <- NULL; rv$last_accept <- NA
updateActionButton(session, "play_pause", label = "▶ Play")
}
isolate(reset_run())
do_step <- function() {
if (rv$done) return(invisible(NULL))
draw <- sample_hull(rv$hull)
x <- draw$x; logg <- draw$logg; logf <- h(x)
rv$tries <- rv$tries + 1
accept <- log(runif(1)) < logf - logg
rv$last_x <- x; rv$last_accept <- accept
if (accept) {
rv$samples <- c(rv$samples, x)
} else {
rv$hull <- build_hull(sort(c(rv$hull$x, x)), rv$lb, rv$ub)
}
if (length(rv$samples) >= rv$target_n) {
rv$done <- TRUE
rv$playing <- FALSE
updateActionButton(session, "play_pause", label = "▶ Play")
}
}
observeEvent(input$l, { reset_run() }, ignoreInit = TRUE)
observeEvent(input$u, { reset_run() }, ignoreInit = TRUE)
observeEvent(input$restart, { reset_run() })
observeEvent(input$n, {
req(input$n)
rv$target_n <- input$n
if (length(rv$samples) < rv$target_n) rv$done <- FALSE
}, ignoreInit = TRUE)
observeEvent(input$play_pause, {
if (rv$done) return()
rv$playing <- !rv$playing
updateActionButton(session, "play_pause",
label = if (rv$playing) "⏸ Pause" else "▶ Play")
})
observeEvent(input$next_step, {
rv$playing <- FALSE
updateActionButton(session, "play_pause", label = "▶ Play")
do_step()
})
### auto-advance one step every 120ms while playing
observe({
if (isTRUE(rv$playing) && !rv$done) {
invalidateLater(120, session)
isolate(do_step())
}
})
output$status <- renderText({
sprintf("%d / %d samples accepted | %d proposals so far | %d tangent points",
length(rv$samples), rv$target_n, rv$tries, length(rv$hull$x))
})
output$panels <- renderPlot({
hull <- rv$hull; lb <- rv$lb; ub <- rv$ub
dl <- smart_xlim(lb, ub)
xg <- seq(dl[1], dl[2], length.out = 400)
hull_log <- hull_val(hull, xg); target_log <- h(xg)
par(mfrow = c(2, 1), mar = c(4, 4, 3, 1))
plot(xg, hull_log, type = "l", lwd = 2, col = "steelblue",
ylim = range(c(hull_log, target_log)), xlab = "x",
ylab = "log density (unnormalized)",
main = sprintf("log target vs. hull (%d tangent points)", length(hull$x)))
lines(xg, target_log, lwd = 2, col = "firebrick")
points(hull$x, hull$hx, pch = 19, col = "gray20", cex = 0.9)
abline(v = hull$z[-c(1, length(hull$z))], col = adjustcolor("gray40", 0.4), lty = 3)
abline(v = c(lb, ub), col = "darkorange", lty = 2)
if (!is.null(rv$last_x)) {
points(rv$last_x, h(rv$last_x), pch = if (rv$last_accept) 19 else 4,
col = if (rv$last_accept) "forestgreen" else "black", cex = 1.6, lwd = 2)
}
legend("bottomleft", bty = "n", cex = 0.85,
lwd = c(2, 2, 1, NA, NA), lty = c(1, 1, 2, NA, NA), pch = c(NA, NA, NA, 19, 4),
col = c("steelblue", "firebrick", "darkorange", "forestgreen", "black"),
legend = c("log hull", "log target", "truncation limits", "last: accepted", "last: rejected"))
if (length(rv$samples) == 0) {
plot(0, 0, type = "n", xlim = dl, ylim = c(0, 1), xlab = "x", ylab = "",
main = "no samples accepted yet")
} else {
### scale the bin count so ~30 bins fall within the *visible* view,
### not the full [lb, ub] window -- otherwise a wide window (e.g.
### [100, 200] zoomed to [100, 102]) puts everything in one giant bin
n_bins <- min(5000, max(30, ceiling(30 * (ub - lb) / (dl[2] - dl[1]))))
breaks <- seq(lb, ub, length.out = n_bins + 1)
hist(rv$samples, breaks = breaks, freq = FALSE, col = "steelblue", border = "white",
xlim = dl, xlab = "x",
main = sprintf("%d / %d samples%s", length(rv$samples), rv$target_n,
if (rv$done) " (done)" else ""))
curve(tnorm_density(x, lb, ub), add = TRUE, lwd = 3, col = "firebrick", n = 300)
legend("topright", bty = "n", lwd = 3, col = "firebrick", legend = "true truncated-normal density")
}
})
}
shinyApp(ui, server)This app accompanies Rejection Sampling in the book.