Skip to content

Adaptive radon centering ​

Counties with many observations can benefit from different random-effect coordinates than counties with few. This example fits a county intercept and floor slope to all 12,573 observations in 386 counties, then compares noncentering, offline selection and online adaptation.

The PosteriorDB model ​

We reproduce radon_all-radon_variable_intercept_slope_noncentered at revision 5545a1dd07ae297c36edecbcd82aa49097b4c385. The model has 772 hierarchical effect cells and 777 scalar parameters, tied with its centered twin for the largest radon model in this PosteriorDB inventory.

text
log_radon[i] ~ Normal(mu_alpha + alpha[county[i]]
                     + (mu_beta + beta[county[i]]) * floor[i], sigma_y)
mu_alpha, mu_beta ~ Normal(0, 10)
sigma_alpha, sigma_beta, sigma_y ~ Normal(0, 1), restricted to positive values
alpha[j] = sigma_alpha * alpha_raw[j]
beta[j] = sigma_beta * beta_raw[j]
alpha_raw[j], beta_raw[j] ~ Normal(0, 1)

The BRM formula uses two named random-effect blocks. Their separate identities keep the county intercept and slope independent and let us specify each half-normal scale prior explicitly.

brm-comparison
The full radon model
julia
using BayesianRegressionModels, Distributions, StanBlocks, JSON, ZipFile

function adaptive_radon_model()
    zip_path = joinpath(dirname(pathof(BayesianRegressionModels)), "..",
                        "research", "radon_centering", "reference", "radon_all.json.zip")
    reader = ZipFile.Reader(zip_path)
    raw = try
        length(reader.files) == 1 || error("radon_all archive must contain one JSON file")
        read(reader.files[1], String)
    finally
        close(reader)
    end
    parsed = JSON.parse(raw)
    N = parsed["N"]
    J = parsed["J"]
    N == 12_573 && J == 386 || error("unexpected radon_all dimensions")
    floor_measure = Float64.(parsed["floor_measure"])
    log_radon = Float64.(parsed["log_radon"])
    county_idx = Int.(parsed["county_idx"])
    (@brm begin
        sigma_y ~ Normal(0, 1; lower=0)
        mu ~ 1 + floor_measure +
              (1 | county_intercept | county_idx) +
              (0 + floor_measure | county_slope | county_idx)
        effect(mu, Intercept) ~ Normal(0, 10)
        effect(mu, floor_measure) ~ Normal(0, 10)
        sd(:, county_intercept) ~ Normal(0, 1)
        sd(:, county_slope) ~ Normal(0, 1)
        log_radon ~ Normal(mu, sigma_y)
    end)((; floor_measure, county_idx, log_radon))
end
julia
BRMI:
  sigma_y ~ Normal(0, 1; lower=0)
  floor_measure: data (eltype=Float64, n=12573)
  county_idx: data (eltype=Int64, n=12573)
  mu ~ 1 + floor_measure + (1 | county_intercept | county_idx) + ((0 + floor_measure) | county_slope | county_idx)
  effect(mu, Intercept) ~ Normal(0, 10)
  effect(mu, floor_measure) ~ Normal(0, 10)
  effect(sd, county_intercept) ~ Normal(0, 1)
  effect(sd, county_slope) ~ Normal(0, 1)
  log_radon ~ Normal(mu, sigma_y)
julia
SBBRMI with data keys = [:county_idx_idx, :floor_measure, :log_radon, :n_county_idx, :n_terms_county_intercept_county_idx, :n_terms_county_slope_county_idx]
configured submodels:
ranef_correlated_draws_generic_configured_1 = Base.merge(BayesianRegressionModels.ranef_correlated_draws_generic, quote
            tau ~ normal(0, 1; n = n_terms, lower = 0.0)
        end)
emitted @slic body:
begin
    b_county_intercept_county_idx ~ ranef_correlated_draws_generic_configured_1(; group_idx = county_idx_idx, n_groups = n_county_idx, n_terms = n_terms_county_intercept_county_idx, lkj_eta = 1.0)
    b_county_slope_county_idx ~ ranef_correlated_draws_generic_configured_1(; group_idx = county_idx_idx, n_groups = n_county_idx, n_terms = n_terms_county_slope_county_idx, lkj_eta = 1.0)
    sigma_y ~ normal(0, 1; lower = 0.0)
    X_mu = hcat(rep_vector(1.0, num_elements(floor_measure)), floor_measure)
    pop_mu ~ _popefs_normal(; X = X_mu, beta_loc = [0, 0], beta_scale = [10, 10])
    r_mu_county_intercept_county_idx = b_county_intercept_county_idx[county_idx_idx, 1]
    r_mu_county_slope_county_idx = floor_measure .* b_county_slope_county_idx[county_idx_idx, 1]
    mu = pop_mu + r_mu_county_intercept_county_idx + r_mu_county_slope_county_idx
    log_radon ~ normal(mu, sigma_y)
end
stan
functions {
matrix hcat(
    vector x,
    vector y
) {
    int n = dims(x)[1];
    if (dims(y)[1] != n) reject("hcat: dim mismatch — `y` dim 1 (= ", dims(y)[1], ") does not match `n` (= ", n, "), inferred from `x` dim 1. `n` sizes: `x` dim 1 (= ", dims(x)[1], "), `y` dim 1 (= ", dims(y)[1], ").");
    return append_col(x, y);
}
vector normal_lpdfs(
    vector obs,
    vector loc,
    real scale
) {
    return jbroadcasted_normal_lpdfs(obs, loc, scale);
}
vector jbroadcasted_normal_lpdfs(
    vector x1,
    vector x2,
    real x3
) {
    int n = dims(x1)[1];
    vector[n] rv;
    for(i in 1:n) {
        rv[i] = normal_lpdfs(broadcasted_getindex(x1, i), broadcasted_getindex(x2, i), x3);
    }
    return rv;
}
real normal_lpdfs(
    real args1,
    real args2,
    real args3
) {
    return normal_lpdf(args1 | args2, args3);
}
real broadcasted_getindex(vector x, int i) {
    return x[i];
}
vector normal_vector_rng(
    int anontok__1,
    vector a,
    real b
) {
    int n = anontok__1;
    if((n == 0)) {
        vector[n] rv;
        return rv;
    } else {
        return to_vector(normal_rng(a, b));
    }
}
}
data {
    int n_terms_county_intercept_county_idx;
    int n_county_idx;
    int n_terms_county_slope_county_idx;
    int floor_measure_n;
    vector[floor_measure_n] floor_measure;
    int county_idx_idx_n;
    array[county_idx_idx_n] int county_idx_idx;
    int log_radon_n;
    vector[log_radon_n] log_radon;
}
transformed data {
    matrix[floor_measure_n, 2] X_mu = hcat(rep_vector(1.0, num_elements(floor_measure)), floor_measure);
    int pop_mu_n_covariates = 2;
}
parameters {
    cholesky_factor_corr[n_terms_county_intercept_county_idx] b_county_intercept_county_idx_L;
    vector<lower=0.0>[n_terms_county_intercept_county_idx] b_county_intercept_county_idx_tau;
    vector[(n_terms_county_intercept_county_idx * n_county_idx)] b_county_intercept_county_idx_z_flat;
    cholesky_factor_corr[n_terms_county_slope_county_idx] b_county_slope_county_idx_L;
    vector<lower=0.0>[n_terms_county_slope_county_idx] b_county_slope_county_idx_tau;
    vector[(n_terms_county_slope_county_idx * n_county_idx)] b_county_slope_county_idx_z_flat;
    real<lower=0.0> sigma_y;
    vector[pop_mu_n_covariates] pop_mu_beta_pop;
}
transformed parameters {
    matrix[n_terms_county_intercept_county_idx, n_county_idx] b_county_intercept_county_idx_z = to_matrix(b_county_intercept_county_idx_z_flat, n_terms_county_intercept_county_idx, n_county_idx);
    matrix[n_county_idx, n_terms_county_intercept_county_idx] b_county_intercept_county_idx = ((
        diag_pre_multiply(b_county_intercept_county_idx_tau, b_county_intercept_county_idx_L) *
        b_county_intercept_county_idx_z
    )');
    matrix[n_terms_county_slope_county_idx, n_county_idx] b_county_slope_county_idx_z = to_matrix(b_county_slope_county_idx_z_flat, n_terms_county_slope_county_idx, n_county_idx);
    matrix[n_county_idx, n_terms_county_slope_county_idx] b_county_slope_county_idx = ((
        diag_pre_multiply(b_county_slope_county_idx_tau, b_county_slope_county_idx_L) *
        b_county_slope_county_idx_z
    )');
    vector[floor_measure_n] pop_mu = (X_mu * pop_mu_beta_pop);
    vector[county_idx_idx_n] r_mu_county_intercept_county_idx = b_county_intercept_county_idx[county_idx_idx, 1];
    vector[county_idx_idx_n] r_mu_county_slope_county_idx = (floor_measure .* b_county_slope_county_idx[county_idx_idx, 1]);
    vector[floor_measure_n] mu = (pop_mu + r_mu_county_intercept_county_idx + r_mu_county_slope_county_idx);
}
model {
    b_county_intercept_county_idx_L ~ lkj_corr_cholesky(1.0);
    b_county_intercept_county_idx_tau ~ normal(0, 1);
    b_county_intercept_county_idx_z_flat ~ std_normal();
    b_county_slope_county_idx_L ~ lkj_corr_cholesky(1.0);
    b_county_slope_county_idx_tau ~ normal(0, 1);
    b_county_slope_county_idx_z_flat ~ std_normal();
    sigma_y ~ normal(0, 1);
    pop_mu_beta_pop ~ normal([0, 0]', [10, 10]');
    log_radon ~ normal(mu, sigma_y);
}
generated quantities {
    vector[log_radon_n] log_radon_likelihood = normal_lpdfs(log_radon, mu, sigma_y);
    vector[log_radon_n] log_radon_gen = normal_vector_rng(log_radon_n, mu, sigma_y);
}
julia
#= line 0 =# Turing.@model(function brm_model(y, X_mu, group_effects_mu_1, group_effects_mu_2)
        sigma_y ~ BayesianRegressionModelsTuringExt._brm_constrained_kernel(Distributions.Normal(0, 1); lower = 0)
        beta_pop ~ Distributions.product_distribution([Distributions.Normal(0, 10), Distributions.Normal(0, 10)])
        eta_mu = X_mu * beta_pop
        group_effect_1 = Base.zeros(Base.length(y))
        group_1_1 ~ DynamicPPL.to_submodel(BRM.turing_group_effect(group_effects_mu_1, (Distributions.Normal(0, 1),), nothing))
        group_effect_1 = group_effect_1 + group_1_1.effect
        group_1_2 ~ DynamicPPL.to_submodel(BRM.turing_group_effect(group_effects_mu_2, (Distributions.Normal(0, 1),), nothing))
        group_effect_1 = group_effect_1 + group_1_2.effect
        eta_mu = eta_mu + group_effect_1
        mu = eta_mu
        begin
            for i = Base.eachindex(y)
                y[i] ~ Distributions.Normal(mu[i], sigma_y)
            end
        end
        (; mu = mu, sigma_y = sigma_y, response = y)
    end)

The tabs show the generated backends. Sampling uses StanBlocks/BridgeStan. An independent audit compares the normalized density and all 777 gradient components with the reference Stan program at six points spanning small and large scales. The maximum absolute density and gradient errors are 1.42e-10 and 8.55e-11.

Fit and check the noncentered model ​

julia
using Random, WarmupHMC

sb = SBBRMI(adaptive_radon_model(); mod=@__MODULE__, total_groups=())
stan_problem = StanBlocks.stan_instantiate(sb.model)
pilot = WarmupHMC.adaptive_warmup_mcmc(
    Xoshiro(1), stan_problem; n_draws=10_000, monitor_ess=true)

Every fit uses one chain and 10,000 retained draws, with ordinary WarmupHMC initialization and adaptation defaults. The same seed and model are used for the selected-partial and online fits.

Each vertical interval predicts one observation: the thin interval contains 90% of replicated values and the thick interval 50%. Observed values are overlaid as points coloured by their recorded floor code. The horizontal position is the original data-row index; observations are not reordered by floor or response, and adjacent intervals are not connected.

The panels show County 19 (11 observations), County 345 (39 observations), County 202 (765 observations). They span the minimum, median and maximum sample sizes among counties with at least five observations at both floor codes 0 and 1. All observations in those counties are shown. The model and the complete predictive table use all 12,573 rows.

The supplied floor covariate contains codes 0, 1, 2, 3 and 9 (with counts 8,299, 3,949, 22, 19 and 284). The formula uses those numeric values as supplied by PosteriorDB; plot labels identify them as codes.

Select one centering per county effect ​

Each county deviation has coordinates u=scz. Its population coefficient remains separate. The following position-loss selection freezes an independent control for every county intercept and slope:

julia
using Enzyme, BridgeStan
using DifferentiationInterface: AutoEnzyme
backend = AutoEnzyme(;
    mode=Enzyme.set_runtime_activity(Enzyme.Reverse),
    function_annotation=Enzyme.Const)
blocks = adaptive_centering_blocks(sb, BridgeStan.param_unc_names(stan_problem.model))
selected = Dict{Int,Float64}()
for block in blocks
    q = pilot.posterior_position
    z = permutedims(q[vec(block.effects), :])
    logscale = vec(q[only(block.log_scales), :])
    choice = select_ranef_centeredness(z, repeat(logscale, 1, size(z, 2));
        candidates=0:0.01:1)
    for (index, c) in zip(vec(block.effects), choice.centeredness)
        selected[index] = c
    end
end
selected_problem = adaptive_centering_problem(sb, stan_problem, backend)
WarmupHMC.restore_reparam_sources!(selected_problem,
    [index => WarmupHMC.PartiallyCentered(selected[index])
     for (index, _) in WarmupHMC.reparam_sources(selected_problem)])
refit = WarmupHMC.adaptive_warmup_mcmc(
    Xoshiro(1), selected_problem;
    n_draws=10_000, monitor_ess=true, nonlinear_adapt=false)

Both losses, post-hoc and online ​

For a zero-mean random effect with log scale ℓ, the centering family is uc=exp⁡(cℓ)z: c=0 is NCP and c=1 is CP. Its transformed effect gradient is gc=exp⁡(−cℓ)gz. We compare two criteria, minimized separately for each effect:

Lposition(c)=log⁡sd(uc)−mean(cℓ),Lgradient(c)=cor(uc,gc).

The second is a signed correlation: an independent Gaussian coordinate has correlation −1 with its log-density gradient. The first uses positions and the Jacobian, without a gradient term.

Both post-hoc arms use the same NCP pilot, select on 0:0.01:1, then fit afresh with the controls fixed. The gradient selector uses the pilot's saved gradients. Both online arms select on the native 0:0.1:1 grid during warmup, using the sampler's trajectory evidence and weights. They have no separate pilot. Controls are frozen for the retained sampling phase.

The online gradient criterion is WarmupHMC's default. The research harness selects the position criterion through the existing internal loss functions; there is currently no public loss-selection keyword. The model, initialization policy, seed and requested draw count are otherwise shared across the arms.

These runs use the active-position transport implementation published in WarmupHMC 6b377cb. When centering changes, the active position and the adaptation sample now represent the same physical points before and after the change.

WarmupHMC returns model coordinates in posterior_position; checkpoints retain sampler coordinates. Export checks their mapping, Jacobian-adjusted density and saved gradients before making figures or scientific summaries.

Sampling efficiency and full workflow cost ​

Each completed arm has one chain, seed 1 and 10,000 retained draws. Every row uses the same scientific quantities: two population coefficients, two group SDs, residual SD, and all 386 county intercept totals and 386 slope totals (777 quantities). Standardized effects are excluded from the minimum. Positive scales may be stored as logs; rank-normalized bulk ESS is invariant under that monotone change.

WHMC methodTotal gradientsSampling efficiencyTotal efficiency
NCP164,9171×1×
CP152,5460.64×0.647×
Post-hoc position317,1861.07×0.52×
Post-hoc gradient319,3851.35×0.651×
Online position154,5390.843×0.841×
Online gradient154,4480.743×0.742×

Both efficiency columns are relative to this study's NCP + WarmupHMC baseline. Sampling efficiency is minimum bulk ESS divided by retained-sampling gradient calls. Total efficiency divides that same minimum ESS by the full workflow's gradient calls. The total includes initialization, all warmup and adaptation, active-state reevaluations, and sampling. For each post-hoc row it also includes the entire NCP pilot; the pilot's ESS is not added to the refit's.

Gradient counts measure target evaluations, a proxy for compute cost rather than a wall-clock speed ratio. Compilation, plotting and independent audits are outside the fitting counts. The complete numerical summaries, including absolute ESS and both denominators, are in the linked result files.

Sampling divergences: NCP: 0; CP: 0; Post-hoc position: 0; Post-hoc gradient: 0; Online position: 0; Online gradient: 0. These are one-chain comparisons, so neither the ranking nor a within-chain split R-hat establishes cross-chain convergence.

The population floor slope limits minimum ESS in every arm. A more favorable local effect geometry does not necessarily improve this global bottleneck. In this run, neither online loss beats NCP in total efficiency.

For each random-effect role, the scatter rows select the minimum post-hoc position-loss centeredness, the value nearest 0.5, and the maximum. Ties use the lowest county index, with distinct counties in the three rows. These same coordinates are used throughout. The selection and full fits include all 386 counties; the selected county labels appear in the figures and exported coordinate table.

Geometry of the fitted coordinates ​

The left column is always the centered visualization baseline, obtained from the NCP pilot. The pilot comparison uses those same draws in CP and NCP. The post-hoc and online panels each show their two newly fitted loss variants in the coordinates actually used by the sampler. CP is the visual reference; NCP remains the efficiency baseline. Axes are independent across panels.

Position and gradient in the displayed coordinates ​

These panels pair each displayed effect coordinate with its own log-density gradient, using 1,000 evenly spaced retained draws. The CP reference transforms the pilot's positions and gradients together. The fitted panels use gradients saved in their actual sampler frame; they do not attach an NCP gradient to a centered position.

Reproduce and inspect ​

The refresh harness contains the driver, both loss selectors, saved-frame audits, export and AoV plotting code. Its results for this study contain the full efficiency denominators, per-quantity ESS and selected controls. The original model directory retains the source specification and independent density/gradient audit. The refresh uses those same model definitions.

Run run.jl radon ncp OUTPUT first, then request cp,posthoc_position,posthoc_gradient,online_position,online_gradient with the same output root. Completed arm directories are immutable. See the harness README for the environment and full commands.