Eight schools and the choice of centering
Eight schools report treatment-effect estimates with substantial uncertainty. A hierarchical model pools those estimates toward a common mean. When the between-school standard deviation is small, centered school effects occupy a narrow region around that mean. Noncentered coordinates separate the scale from standardized school effects and can make this region easier to sample.
This example compares a manually fully centered fit with a noncentered fit, then uses the same model to demonstrate offline and online centering selection.
The PosteriorDB model
We reproduce PosteriorDB's eight-schools posterior, including its data and priors:
mu ~ Normal(0, 5)
tau ~ Cauchy(0, 5), restricted to positive values
z[j] ~ Normal(0, 1)
theta[j] = mu + tau*z[j]
y[j] ~ Normal(theta[j], sigma[j])The positive restriction on tau makes its prior half-Cauchy. The sampling statement and support match the reference Stan program. Here, sigma contains the known standard errors of the reported estimates.
Eight schools with PosteriorDB priorsusing BayesianRegressionModels, Distributions, StanBlocks
function eight_schools_model()
(@brm begin
theta ~ 1 + (1 | eight_schools | school)
effect(theta, Intercept) ~ Normal(0, 5)
sd(:, eight_schools) ~ Cauchy(0, 5)
y ~ Normal(theta, sigma)
end)((; school=1:8,
y=Float64[28, 8, -3, 7, -1, 1, 18, 12],
sigma=Float64[15, 10, 16, 11, 9, 11, 10, 18]))
endBRMI:
school: data (eltype=Int64, n=8)
theta ~ 1 + (1 | eight_schools | school)
effect(theta, Intercept) ~ Normal(0, 5)
effect(sd, eight_schools) ~ Cauchy(0, 5)
sigma: data (eltype=Float64, n=8)
y ~ Normal(theta, sigma)SBBRMI with data keys = [:n_school, :n_terms_eight_schools_school, :school_idx, :sigma, :y]
configured submodels:
ranef_correlated_draws_generic_configured_1 = Base.merge(BayesianRegressionModels.ranef_correlated_draws_generic, quote
tau ~ cauchy(0, 5; n = n_terms, lower = 0.0)
end)
emitted @slic body:
begin
b_eight_schools_school ~ ranef_correlated_draws_generic_configured_1(; group_idx = school_idx, n_groups = n_school, n_terms = n_terms_eight_schools_school, lkj_eta = 1.0)
X_theta = hcat(rep_vector(1.0, num_elements(school_idx)))
pop_theta ~ _popefs_normal(; X = X_theta, beta_loc = [0], beta_scale = [5])
r_theta_eight_schools_school = b_eight_schools_school[school_idx, 1]
theta = pop_theta + r_theta_eight_schools_school
y ~ normal(theta, sigma)
endfunctions {
matrix hcat(vector x) {
int n = dims(x)[1];
return to_matrix(x, n, 1);
}
vector normal_lpdfs(
vector obs,
vector loc,
vector scale
) {
return jbroadcasted_normal_lpdfs(obs, loc, scale);
}
vector jbroadcasted_normal_lpdfs(
vector x1,
vector x2,
vector 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), broadcasted_getindex(x3, i));
}
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,
vector 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_eight_schools_school;
int n_school;
int school_idx_n;
array[school_idx_n] int school_idx;
int y_n;
vector[y_n] y;
int sigma_n;
vector[sigma_n] sigma;
}
transformed data {
matrix[num_elements(school_idx), 1] X_theta = hcat(rep_vector(1.0, num_elements(school_idx)));
int pop_theta_n_covariates = 1;
}
parameters {
cholesky_factor_corr[n_terms_eight_schools_school] b_eight_schools_school_L;
vector<lower=0.0>[n_terms_eight_schools_school] b_eight_schools_school_tau;
vector[(n_terms_eight_schools_school * n_school)] b_eight_schools_school_z_flat;
vector[pop_theta_n_covariates] pop_theta_beta_pop;
}
transformed parameters {
matrix[n_terms_eight_schools_school, n_school] b_eight_schools_school_z = to_matrix(b_eight_schools_school_z_flat, n_terms_eight_schools_school, n_school);
matrix[n_school, n_terms_eight_schools_school] b_eight_schools_school = ((diag_pre_multiply(b_eight_schools_school_tau, b_eight_schools_school_L) * b_eight_schools_school_z)');
vector[num_elements(school_idx)] pop_theta = (X_theta * pop_theta_beta_pop);
vector[school_idx_n] r_theta_eight_schools_school = b_eight_schools_school[school_idx, 1];
vector[num_elements(school_idx)] theta = (pop_theta + r_theta_eight_schools_school);
}
model {
b_eight_schools_school_L ~ lkj_corr_cholesky(1.0);
b_eight_schools_school_tau ~ cauchy(0, 5);
b_eight_schools_school_z_flat ~ std_normal();
pop_theta_beta_pop ~ normal([0]', [5]');
y ~ normal(theta, sigma);
}
generated quantities {
vector[y_n] y_likelihood = normal_lpdfs(y, theta, sigma);
vector[y_n] y_gen = normal_vector_rng(y_n, theta, sigma);
}#= line 0 =# Turing.@model(function brm_model(y, X_theta, group_effects_theta_1, sigma)
beta_pop ~ Distributions.product_distribution([Distributions.Normal(0, 5)])
eta_theta = X_theta * beta_pop
group_effect_1 = Base.zeros(Base.length(y))
group_1_1 ~ DynamicPPL.to_submodel(BRM.turing_group_effect(group_effects_theta_1, (Distributions.Cauchy(0, 5),), nothing))
group_effect_1 = group_effect_1 + group_1_1.effect
eta_theta = eta_theta + group_effect_1
theta = eta_theta
begin
for i = Base.eachindex(y)
y[i] ~ Distributions.Normal(theta[i], sigma[i])
end
end
(; theta = theta, response = y)
end)The tabs show generated backends. The fits below use StanBlocks/BridgeStan and WarmupHMC. Independent comparisons with PosteriorDB's centered Stan program verify the density and all ten gradient components, including the coordinate Jacobian. The independent audit is retained with the source model.
Choose fully centered coordinates manually
The conventional BRM model (total_groups=()) samples standardized effects z[j]. Full centering of the random effects instead samples u[j] = tau*z[j], so u[j] ~ Normal(0, tau) and theta[j] = mu + u[j]. Thus u is the school's deviation from the population mean; theta is its treatment effect.
To select this parameterization when compiling the model, use the grouping factor's name:
model = eight_schools_model()
centered_sb = SBBRMI(model;
mod=@__MODULE__, centered_groups=[:school])We can also choose the same endpoint through the centering controls. This lets all six fits below share one compiled noncentered density:
using Random, Enzyme, WarmupHMC, BridgeStan
using DifferentiationInterface: AutoEnzyme
sb = SBBRMI(model; mod=@__MODULE__, total_groups=())
problem = StanBlocks.stan_instantiate(sb.model)
backend = AutoEnzyme(;
mode=Enzyme.set_runtime_activity(Enzyme.Reverse),
function_annotation=Enzyme.Const)
function fixed_centering(sb, problem, backend, coefficients)
target = adaptive_centering_problem(sb, problem, backend)
sources = WarmupHMC.reparam_sources(target)
length(sources) == length(coefficients) || error("one coefficient per school is required")
WarmupHMC.restore_reparam_sources!(target,
[index => WarmupHMC.PartiallyCentered(Float64(c))
for ((index, _), c) in zip(sources, coefficients)])
target
end
fully_centered = fixed_centering(sb, problem, backend, ones(8))
centered_fit = WarmupHMC.adaptive_warmup_mcmc(
Xoshiro(1), fully_centered;
n_draws=10_000, monitor_ess=true, nonlinear_adapt=false)
noncentered_fit = WarmupHMC.adaptive_warmup_mcmc(
Xoshiro(1), problem; n_draws=10_000, monitor_ess=true)Every coefficient is 1.0; nonlinear_adapt=false keeps that choice fixed. At 0.0 the corresponding coordinate is fully noncentered. Intermediate values sample u[j] = tau^c[j]*z[j].
WarmupHMC returns model coordinates in posterior_position, including when centering is fixed. Its checkpoints retain sampler coordinates. The checkpoint-to-model mapping is checked independently before summarizing fits.
Posterior effects and predictive checks
Thin intervals contain 90% of draws and thick intervals 50%. Each school is a category with its own interval. Treatment effects are
The predictive check uses the NCP fit and includes the known standard error of each reported estimate. Red points are the observations in their original school order. There are no ribbons between categorical schools.
Select centering from a pilot
An offline rule evaluates each school on a grid c = 0:0.01:1:
Compute it using the noncentered pilot's model coordinates, then hold the selected coefficients fixed during a new fit:
blocks = adaptive_centering_blocks(sb, BridgeStan.param_unc_names(problem.model))
block = only(blocks)
q = noncentered_fit.posterior_position
z = permutedims(q[vec(block.effects), :])
logtau = vec(q[only(block.log_scales), :])
selection = select_ranef_centeredness(z, repeat(logtau, 1, 8);
candidates=0:0.01:1)
selected_problem = fixed_centering(sb, problem, backend, selection.centeredness)
partial_fit = 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 second is a signed correlation: an independent Gaussian coordinate has correlation
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.
Both criteria favor coordinates close to noncentering in this weakly informed hierarchy.
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: population mean, group SD and eight school treatment effects (10 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 method | Total gradients | Sampling efficiency | Total efficiency |
|---|---|---|---|
| NCP | 70,276 | 1× | 1× |
| CP | 174,583 | 0.0139× | 0.0139× |
| Post-hoc position | 141,296 | 1.1× | 0.553× |
| Post-hoc gradient | 158,294 | 0.651× | 0.364× |
| Online position | 86,499 | 0.895× | 0.898× |
| Online gradient | 104,806 | 0.614× | 0.618× |
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: 1; CP: 39; Post-hoc position: 0; Post-hoc gradient: 0; Online position: 2; Online gradient: 0. These are one-chain comparisons, so neither the ranking nor a within-chain split R-hat establishes cross-chain convergence.
The position-loss post-hoc refit modestly improves sampling efficiency here, but its pilot makes total efficiency lower than NCP. Full centering has 39 divergences and poor efficiency; its intervals require that qualification.
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 eight 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.








