Skip to content

Pupil: hierarchical residual SD ​

Result ​

Automatic BRM totals work for both the mean and residual-scale predictors in this pupil model. All 15 matched arms completed. With WHMC held fixed, online position adaptation of BRM totals gives 2.36× the total-cost efficiency of brms S2Z auto in this pilot. Against ordinary brms NCP with native Stan, the gain is 400× per sampling gradient and 442× per total gradient.

The results support combining marginalization with centering adaptation. Fixed NCP is weak for both totals and S2Z here. These are single-chain exploratory results; the precise ranking needs replication.

Model and data ​

This follows Aki's post-4 pupil model, with the requested simplification to independent mean random effects:

r
bf(p_size ~ load + (load || subj), sigma ~ (1 | subj))

There are 2,228 observations and 20 subjects, with numeric load 0–5. Observation order is preserved. Residual SD has a subject-specific hierarchical intercept, unlike the numeric-subject scale regression in the previous pupil example.

Writing x̄ for average load:

text
y[n] ~ Normal(beta0 + beta1*(load[n]-x̄) + a[j] + b[j]*load[n],
              exp(gamma0 + v[j]))

The deviations a, b and v are independent zero-mean Gaussians with SDs tau_a, tau_b and tau_v. Priors are taken from the generated brms target:

  • beta0: Student-t(3, 5651.9, 2026.1), at average load;

  • beta1: flat, as in the source model;

  • gamma0: Student-t(3, 0, 2.5);

  • all three group SDs: half-Student-t(3, 0, 2026.1).

The broad prior on tau_v is present in brms's generated code and is retained across every arm. Independent mean random effects are an explicit simplification of the forum model. Every row targets this same simplified posterior.

What BRM does automatically ​

BRM lowers the ordinary formula and prior declarations into subject totals:

text
A[j] = beta0 − x̄*beta1 + a[j]
B[j] = beta1 + b[j]
C[j] = gamma0 + v[j]
y[n] ~ Normal(A[j] + B[j]*load[n], exp(C[j]))

It integrates out the three population coefficients. Each Student-t population prior is represented exactly by a Gaussian conditional on a Gamma(3/2, rate=3/2) precision multiplier. There are two such multipliers. Ordinary brms samples 66 parameters; totals and S2Z each sample 65.

For each total block T, the prior evaluated during sampling is the Gaussian integral over its population coefficient beta. Conditional precision is Q = prior_precision + J*A' * diag(1/tau²) * A; conditional mean is Q \ (prior_precision*prior_location + A' * diag(1/tau²) * sum(T)). The small design-basis matrix A accounts for centering the population load predictor. Conditional recovery uses this same distribution. Evaluation is linear in the number of subjects for a fixed number of coefficients.

BRM discovers the two total blocks, scales and centering controls and connects them to WHMC. The formula and priors determine the compiled posterior and centering map. The linked harness uses the public adaptive_centering_problem, select_total_centeredness and recover_population_draws interfaces.

CP uses model-scale totals. NCP scales each total about its prespecified location by its group SD; the resulting marginal prior remains correlated. Post-hoc adaptation chooses one fixed centering control per total from an NCP pilot. Online adaptation changes those controls during WHMC warmup. Both position-only and position–gradient losses are included. For this experiment the loss switch is selected locally in the driver; package files are unchanged.

brms S2Z uses orthonormal contrast coordinates and integrates out population means with the corresponding exact prior adjustment. Its auto row uses the pinned branch's Pathfinder precursor to select fixed centering weights; WHMC reuses those weights and is charged the precursor's cost. S2Z auto here is not online WHMC adaptation.

BRM model and sampling interface ​

The authoring pane below reads the exact declarations used by the sampling harness. The backend panes are regenerated during the docs build. The measurements use the StanBlocks/BridgeStan backend and WHMC; the Turing pane is a generated-model comparison only.

brm-comparison
Pupil means and hierarchical residual SD
julia
using BayesianRegressionModels, Distributions, JSON, DelimitedFiles

function pupil_hierarchical_brm_model()
    reference = JSON.parsefile(joinpath(pkgdir(BayesianRegressionModels),
        "research", "pupil_scale_totals", "reference", "ordinary_ncp.json"))
    data = (;p_size=Float64.(reference["Y"]), load=Float64.(reference["Z_1_2"]),
             subj=Int.(reference["J_1"]))
    
    builder =  @brm begin
        mu ~ 1 + center(load) + (1 | mean_intercept | subj) + (0 + load | mean_slope | subj)
        logsigma ~ 1 + (1 | scale_intercept | subj)
        effect(mu, Intercept) ~ LocationScale(5651.9,2026.1,TDist(3))
        effect(mu, center_load) ~ Flat()
        effect(logsigma, Intercept) ~ LocationScale(0.,2.5,TDist(3))
        sd(:, mean_intercept) ~ LocationScale(0.,2026.1,TDist(3))
        sd(:, mean_slope) ~ LocationScale(0.,2026.1,TDist(3))
        sd(:, scale_intercept) ~ LocationScale(0.,2026.1,TDist(3))
        p_size ~ Normal(mu,exp(logsigma))
    end
    
    builder(data)
end
julia
BRMI:
  load: data (eltype=Float64, n=2228)
  subj: data (eltype=Int64, n=2228)
  mu ~ 1 + center(load) + (1 | mean_intercept | subj) + ((0 + load) | mean_slope | subj)
  logsigma ~ 1 + (1 | scale_intercept | subj)
  effect(mu, Intercept) ~ AffineDistribution(5651.9, 2026.1, TDist(3))
  effect(mu, center_load) ~ Flat()
  effect(logsigma, Intercept) ~ AffineDistribution(0.0, 2.5, TDist(3))
  effect(sd, mean_intercept) ~ AffineDistribution(0.0, 2026.1, TDist(3))
  effect(sd, mean_slope) ~ AffineDistribution(0.0, 2026.1, TDist(3))
  effect(sd, scale_intercept) ~ AffineDistribution(0.0, 2026.1, TDist(3))
  p_size ~ Normal(mu, exp(logsigma))
julia
SBBRMI with data keys = [:load, :p_size, :subj, :total_A_logsigma, :total_A_mu, :total_group_logsigma, :total_group_mu, :total_location_logsigma, :total_location_mu, :total_mixture_shape_logsigma, :total_mixture_shape_mu, :total_ng_logsigma, :total_ng_mu, :total_nk_logsigma, :total_nk_mu, :total_nm_logsigma, :total_nm_mu, :total_np_logsigma, :total_np_mu, :total_precision_logsigma, :total_precision_mu]
configured submodels:
_brm_total_scales_configured_1 = Base.merge(BayesianRegressionModels._brm_total_scales, quote
            tau ~ (ValueFamily(brm_vector_prior_eecc99814bef47ac))(3.0, 0.0, 2026.1, 3.0, 0.0, 2026.1; n = 2)
        end)
_brm_total_scales_configured_2 = Base.merge(BayesianRegressionModels._brm_total_scales, quote
            tau ~ (ValueFamily(brm_vector_prior_b28d90387aac6d90))(3.0, 0.0, 2026.1; n = 1)
        end)
emitted @slic body:
begin
    total_scale_mu ~ _brm_total_scales_configured_1(; n = total_nk_mu)
    total_mixture_mu::vector[total_nm_mu] ~ gamma(total_mixture_shape_mu, total_mixture_shape_mu; lower = 0.0)
    total_conditional_precision_mu = [total_precision_mu[1] * total_mixture_mu[1], total_precision_mu[2]]
    total_mu::matrix[total_ng_mu, total_nk_mu] ~ brm_total(total_scale_mu, total_A_mu, total_location_mu, total_conditional_precision_mu)
    population_mu = brm_total_recover_rng(total_mu, total_scale_mu, total_A_mu, total_location_mu, total_conditional_precision_mu)
    deviation_mu = brm_total_deviations(total_mu, total_A_mu * population_mu)
    total_Z_mu = hcat(rep_vector(1.0, num_elements(total_group_mu)), load)
    mu = rows_dot_product(total_mu[total_group_mu, :], total_Z_mu)
    total_scale_logsigma ~ _brm_total_scales_configured_2(; n = total_nk_logsigma)
    total_mixture_logsigma::vector[total_nm_logsigma] ~ gamma(total_mixture_shape_logsigma, total_mixture_shape_logsigma; lower = 0.0)
    total_conditional_precision_logsigma = [total_precision_logsigma[1] * total_mixture_logsigma[1]]
    total_logsigma::matrix[total_ng_logsigma, total_nk_logsigma] ~ brm_total(total_scale_logsigma, total_A_logsigma, total_location_logsigma, total_conditional_precision_logsigma)
    population_logsigma = brm_total_recover_rng(total_logsigma, total_scale_logsigma, total_A_logsigma, total_location_logsigma, total_conditional_precision_logsigma)
    deviation_logsigma = brm_total_deviations(total_logsigma, total_A_logsigma * population_logsigma)
    total_Z_logsigma = hcat(rep_vector(1.0, num_elements(total_group_logsigma)))
    logsigma = rows_dot_product(total_logsigma[total_group_logsigma, :], total_Z_logsigma)
    p_size ~ normal(mu, (exp)(logsigma))
end
stan
functions {
// value UDF brm_vector_prior_eecc99814bef47ac_lpdf
real brm_vector_prior_eecc99814bef47ac_lpdf(
    vector x,
    real arg_1,
    real arg_2,
    real arg_3,
    real arg_4,
    real arg_5,
    real arg_6
) {
    if((x[1] < 0.0)) {
        return negative_infinity();
    }
    if((x[2] < 0.0)) {
        return negative_infinity();
    }
    return (student_t_lpdf(x[1] | arg_1, arg_2, arg_3) + student_t_lpdf(x[2] | arg_4, arg_5, arg_6));
}
real brm_total_lpdf(
    matrix total,
    vector tau,
    matrix A,
    vector location,
    vector precision
) {
    int j = dims(total)[1];
    int k = dims(total)[2];
    int p = dims(A)[2];
    if (dims(tau)[1] != k) reject("brm_total_lpdf: dim mismatch — `tau` dim 1 (= ", dims(tau)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(A)[1] != k) reject("brm_total_lpdf: dim mismatch — `A` dim 1 (= ", dims(A)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(location)[1] != p) reject("brm_total_lpdf: dim mismatch — `location` dim 1 (= ", dims(location)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], ").");
    if (dims(precision)[1] != p) reject("brm_total_lpdf: dim mismatch — `precision` dim 1 (= ", dims(precision)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], ").");
    matrix[dims(precision)[1], dims(precision)[1]] Q = brm_total_precision(tau, A, precision, j);
    vector[dims(total)[2]] average = brm_total_mean(total);
    vector[dims(precision)[1]] beta = brm_total_conditional_mean(total, tau, A, location, precision, Q);
    vector[dims(total)[2]] residual = (average - (A * beta));
    real quadratic = 0.0;
    real lp = ((-0.5 * ((j * k) - p) * 1.8378770664093453) - (j * sum(log(tau))));
    for(c in 1:k) {
        quadratic += ((j * square(residual[c])) / square(tau[c]));
        for(g in 1:j) {
            quadratic += (square((total[g, c] - average[c])) / square(tau[c]));
        }
    }
    for(a in 1:p) {
        if((precision[a] > 0.0)) {
            lp += (0.5 * (log(precision[a]) - 1.8378770664093453));
            quadratic += (precision[a] * square((beta[a] - location[a])));
        }
    }
    return (lp - (0.5 * (log_determinant(Q) + quadratic)));
}
matrix brm_total_precision(
    vector tau,
    matrix A,
    vector precision,
    int n_groups
) {
    int k = dims(tau)[1];
    int p = dims(A)[2];
    if (dims(A)[1] != k) reject("brm_total_precision: dim mismatch — `A` dim 1 (= ", dims(A)[1], ") does not match `k` (= ", k, "), inferred from `tau` dim 1. `k` sizes: `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(precision)[1] != p) reject("brm_total_precision: dim mismatch — `precision` dim 1 (= ", dims(precision)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `precision` dim 1 (= ", dims(precision)[1], ").");
    matrix[dims(precision)[1], dims(precision)[1]] out = diag_matrix(precision);
    for(a in 1:p) {
        for(b in 1:p) {
            for(c in 1:k) {
                out[a, b] += ((n_groups * A[c, a] * A[c, b]) / square(tau[c]));
            }
        }
    }
    return out;
}
vector brm_total_mean(
    matrix total
) {
    int j = dims(total)[1];
    int k = dims(total)[2];
    vector[k] out = rep_vector(0.0, k);
    for(c in 1:k) {
        out[c] = (sum(total[:, c]) / j);
    }
    return out;
}
vector brm_total_conditional_mean(
    matrix total,
    vector tau,
    matrix A,
    vector location,
    vector precision,
    matrix Q
) {
    int j = dims(total)[1];
    int k = dims(total)[2];
    int p = dims(A)[2];
    if (dims(tau)[1] != k) reject("brm_total_conditional_mean: dim mismatch — `tau` dim 1 (= ", dims(tau)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(A)[1] != k) reject("brm_total_conditional_mean: dim mismatch — `A` dim 1 (= ", dims(A)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(location)[1] != p) reject("brm_total_conditional_mean: dim mismatch — `location` dim 1 (= ", dims(location)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], "), `Q` dim 1 (= ", dims(Q)[1], "), `Q` dim 2 (= ", dims(Q)[2], ").");
    if (dims(precision)[1] != p) reject("brm_total_conditional_mean: dim mismatch — `precision` dim 1 (= ", dims(precision)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], "), `Q` dim 1 (= ", dims(Q)[1], "), `Q` dim 2 (= ", dims(Q)[2], ").");
    if (dims(Q)[1] != p) reject("brm_total_conditional_mean: dim mismatch — `Q` dim 1 (= ", dims(Q)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], "), `Q` dim 1 (= ", dims(Q)[1], "), `Q` dim 2 (= ", dims(Q)[2], ").");
    if (dims(Q)[2] != p) reject("brm_total_conditional_mean: dim mismatch — `Q` dim 2 (= ", dims(Q)[2], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], "), `Q` dim 1 (= ", dims(Q)[1], "), `Q` dim 2 (= ", dims(Q)[2], ").");
    vector[dims(total)[2]] average = brm_total_mean(total);
    vector[dims(location)[1]] natural = (precision .* location);
    for(a in 1:p) {
        for(c in 1:k) {
            natural[a] += ((j * A[c, a] * average[c]) / square(tau[c]));
        }
    }
    return mdivide_left_spd(Q, natural);
}
vector brm_total_recover_rng(
    matrix total,
    vector tau,
    matrix A,
    vector location,
    vector precision
) {
    int j = dims(total)[1];
    int k = dims(total)[2];
    int p = dims(A)[2];
    if (dims(tau)[1] != k) reject("brm_total_recover_rng: dim mismatch — `tau` dim 1 (= ", dims(tau)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(A)[1] != k) reject("brm_total_recover_rng: dim mismatch — `A` dim 1 (= ", dims(A)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `tau` dim 1 (= ", dims(tau)[1], "), `A` dim 1 (= ", dims(A)[1], ").");
    if (dims(location)[1] != p) reject("brm_total_recover_rng: dim mismatch — `location` dim 1 (= ", dims(location)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], ").");
    if (dims(precision)[1] != p) reject("brm_total_recover_rng: dim mismatch — `precision` dim 1 (= ", dims(precision)[1], ") does not match `p` (= ", p, "), inferred from `A` dim 2. `p` sizes: `A` dim 2 (= ", dims(A)[2], "), `location` dim 1 (= ", dims(location)[1], "), `precision` dim 1 (= ", dims(precision)[1], ").");
    matrix[dims(precision)[1], dims(precision)[1]] Q = brm_total_precision(tau, A, precision, j);
    vector[dims(precision)[1]] beta = brm_total_conditional_mean(total, tau, A, location, precision, Q);
    return multi_normal_rng(beta, inverse_spd(Q));
}
matrix brm_total_deviations(
    matrix total,
    vector mu
) {
    int j = dims(total)[1];
    int k = dims(total)[2];
    if (dims(mu)[1] != k) reject("brm_total_deviations: dim mismatch — `mu` dim 1 (= ", dims(mu)[1], ") does not match `k` (= ", k, "), inferred from `total` dim 2. `k` sizes: `total` dim 2 (= ", dims(total)[2], "), `mu` dim 1 (= ", dims(mu)[1], ").");
    matrix[dims(total)[1], dims(total)[2]] out = total;
    for(c in 1:k) {
        out[:, c] = (total[:, c] - rep_vector(mu[c], j));
    }
    return out;
}
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);
}
// value UDF brm_vector_prior_b28d90387aac6d90_lpdf
real brm_vector_prior_b28d90387aac6d90_lpdf(
    vector x,
    real arg_1,
    real arg_2,
    real arg_3
) {
    if((x[1] < 0.0)) {
        return negative_infinity();
    }
    return student_t_lpdf(x[1] | arg_1, arg_2, arg_3);
}
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 total_nm_mu;
    int total_mixture_shape_mu_n;
    vector[total_mixture_shape_mu_n] total_mixture_shape_mu;
    int total_precision_mu_n;
    vector[total_precision_mu_n] total_precision_mu;
    int total_ng_mu;
    int total_nk_mu;
    int total_A_mu_m;
    int total_A_mu_n;
    matrix[total_A_mu_m, total_A_mu_n] total_A_mu;
    int total_location_mu_n;
    vector[total_location_mu_n] total_location_mu;
    int load_n;
    int total_group_mu_n;
    array[total_group_mu_n] int total_group_mu;
    vector[load_n] load;
    int total_nm_logsigma;
    int total_mixture_shape_logsigma_n;
    vector[total_mixture_shape_logsigma_n] total_mixture_shape_logsigma;
    int total_precision_logsigma_n;
    vector[total_precision_logsigma_n] total_precision_logsigma;
    int total_ng_logsigma;
    int total_nk_logsigma;
    int total_A_logsigma_m;
    int total_A_logsigma_n;
    matrix[total_A_logsigma_m, total_A_logsigma_n] total_A_logsigma;
    int total_location_logsigma_n;
    vector[total_location_logsigma_n] total_location_logsigma;
    int total_group_logsigma_n;
    array[total_group_logsigma_n] int total_group_logsigma;
    int p_size_n;
    vector[p_size_n] p_size;
}
transformed data {
    matrix[load_n, 2] total_Z_mu = hcat(rep_vector(1.0, num_elements(total_group_mu)), load);
    matrix[num_elements(total_group_logsigma), 1] total_Z_logsigma = hcat(rep_vector(1.0, num_elements(total_group_logsigma)));
}
parameters {
    vector<lower=0.0>[2] total_scale_mu_tau;
    vector<lower=0.0>[total_nm_mu] total_mixture_mu;
    matrix[total_ng_mu, total_nk_mu] total_mu;
    vector<lower=0.0>[1] total_scale_logsigma_tau;
    vector<lower=0.0>[total_nm_logsigma] total_mixture_logsigma;
    matrix[total_ng_logsigma, total_nk_logsigma] total_logsigma;
}
transformed parameters {
    vector<lower=0.0>[2] total_scale_mu = total_scale_mu_tau;
    vector[2] total_conditional_precision_mu = [(total_precision_mu[1] * total_mixture_mu[1]), total_precision_mu[2]]';
    vector[load_n] mu = rows_dot_product(total_mu[total_group_mu, :], total_Z_mu);
    vector<lower=0.0>[1] total_scale_logsigma = total_scale_logsigma_tau;
    vector[1] total_conditional_precision_logsigma = [(total_precision_logsigma[1] * total_mixture_logsigma[1])]';
    vector[num_elements(total_group_logsigma)] logsigma = rows_dot_product(total_logsigma[total_group_logsigma, :], total_Z_logsigma);
}
model {
    total_scale_mu_tau ~ brm_vector_prior_eecc99814bef47ac(3.0, 0.0, 2026.1, 3.0, 0.0, 2026.1);
    total_mixture_mu ~ gamma(total_mixture_shape_mu, total_mixture_shape_mu);
    total_mu ~ brm_total(total_scale_mu, total_A_mu, total_location_mu, total_conditional_precision_mu);
    total_scale_logsigma_tau ~ brm_vector_prior_b28d90387aac6d90(3.0, 0.0, 2026.1);
    total_mixture_logsigma ~ gamma(total_mixture_shape_logsigma, total_mixture_shape_logsigma);
    total_logsigma ~ brm_total(
        total_scale_logsigma,
        total_A_logsigma,
        total_location_logsigma,
        total_conditional_precision_logsigma
    );
    p_size ~ normal(mu, exp(logsigma));
}
generated quantities {
    vector[2] population_mu = brm_total_recover_rng(
        total_mu,
        total_scale_mu,
        total_A_mu,
        total_location_mu,
        total_conditional_precision_mu
    );
    matrix[total_ng_mu, total_A_mu_m] deviation_mu = brm_total_deviations(total_mu, (total_A_mu * population_mu));
    vector[1] population_logsigma = brm_total_recover_rng(
        total_logsigma,
        total_scale_logsigma,
        total_A_logsigma,
        total_location_logsigma,
        total_conditional_precision_logsigma
    );
    matrix[total_ng_logsigma, total_A_logsigma_m] deviation_logsigma = brm_total_deviations(total_logsigma, (total_A_logsigma * population_logsigma));
    vector[p_size_n] p_size_likelihood = normal_lpdfs(p_size, mu, exp(logsigma));
    vector[p_size_n] p_size_gen = normal_vector_rng(p_size_n, mu, exp(logsigma));
}
julia
#= line 0 =# Turing.@model(function brm_model(y, callable_1, X_mu, group_effects_mu_1, callable_2, group_effects_mu_2, callable_3, callable_4, X_logsigma, group_effects_logsigma_1, callable_5)
        beta_pop ~ Distributions.product_distribution([callable_1(5651.9, 2026.1, Distributions.TDist(3)), BayesianRegressionModels.Flat()])
        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, (callable_2(0.0, 2026.1, Distributions.TDist(3)),), 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, (callable_3(0.0, 2026.1, Distributions.TDist(3)),), nothing))
        group_effect_1 = group_effect_1 + group_1_2.effect
        eta_mu = eta_mu + group_effect_1
        mu = eta_mu
        beta_pop_logsigma ~ Distributions.product_distribution([callable_4(0.0, 2.5, Distributions.TDist(3))])
        eta_logsigma = X_logsigma * beta_pop_logsigma
        group_effect_2 = Base.zeros(Base.length(y))
        group_2_1 ~ DynamicPPL.to_submodel(BRM.turing_group_effect(group_effects_logsigma_1, (callable_5(0.0, 2026.1, Distributions.TDist(3)),), nothing))
        group_effect_2 = group_effect_2 + group_2_1.effect
        eta_logsigma = eta_logsigma + group_effect_2
        logsigma = eta_logsigma
        begin
            for i = Base.eachindex(y)
                y[i] ~ Distributions.Normal(mu[i], Base.exp(logsigma[i]))
            end
        end
        (; mu = mu, logsigma = logsigma, response = y)
    end)
julia
using StanBlocks, BridgeStan, WarmupHMC, Enzyme, Random
using DifferentiationInterface: AutoEnzyme

sb = SBBRMI(pupil_hierarchical_brm_model(); mod=@__MODULE__)
problem = StanBlocks.stan_instantiate(sb.model)
names = BridgeStan.param_unc_names(problem.model)
adaptive = adaptive_centering_problem(sb, problem, AutoEnzyme(); centeredness=0.0)
fit = adaptive_warmup_mcmc(Xoshiro(1), adaptive;
    n_draws=2000, monitor_ess=true, nonlinear_adapt=true)

# Returned positions are in the compiled model frame.
draws = permutedims(fit.posterior_position)
recovered = recover_population_draws(sb, draws, names; rng=Xoshiro(404))

For a fixed endpoint, set centeredness=1.0 (CP) or 0.0 (NCP), and set nonlinear_adapt=false. For a post-hoc position choice, call select_total_centeredness(sb, pilot_draws, names; criterion=:position) and pass its centeredness to a fresh problem and fit. The gradient criterion additionally needs the pilot's gradients in the compiled model frame. The harness transports saved checkpoint gradients into that frame before selection.

One scientific-QOI comparison ​

Every row measures the same 66 quantities: beta0, beta1, gamma0; the three group SDs; and each subject's total intercept A, total slope B, and residual SD exp(C). Marginalized methods recover the population coefficients conditionally. Standardized deviations, raw sampler coordinates and mixture variables are excluded.

Efficiencies are minimum bulk ESS divided by the indicated gradient count, relative to ordinary brms NCP + native Stan. Total gradients include initialization, warmup and centering pilots. Post-hoc totals include their 363,592-gradient NCP pilot in each standalone workflow; S2Z auto includes its 1,786-gradient precursor. The reused native S2Z fit itself is not charged to WHMC.

MethodTotal gradientsSampling efficiencyTotal efficiency
brms NCP · Native Stan2,122,4821×1×
brms NCP · WHMC511,2831.9×2.46×
brms CP · WHMC182,8901×1.14×
brms S2Z CP · Native Stan245,020160×59.4×
brms S2Z CP · WHMC104,514152×148×
brms S2Z NCP · Native Stan1,442,6531.19×1.19×
brms S2Z NCP · WHMC187,7803.44×4.35×
brms S2Z auto · Native Stan171,127218×105×
brms S2Z auto · WHMC97,516209×188×
BRM total NCP · WHMC363,5920.924×1.23×
BRM total CP · WHMC291,873229×37×
BRM total post-hoc position · WHMC413,396182×25×
BRM total post-hoc gradient · WHMC410,621408×42.3×
BRM total online position · WHMC38,606400×442×
BRM total online gradient · WHMC39,224338×376×

Choosing the centering matters greatly. Fixed total CP and the two post-hoc choices have good sampling efficiency, but their adaptation/pilot costs reduce total efficiency. Online selection obtains useful geometry without first paying for the expensive NCP pilot. Ordinary brms CP still mixes the shared population intercept slowly, despite improved mixing of subject totals.

Saved-draw geometry ​

Both figures reuse the same 2,000 draws across their three columns. CP is the visualization baseline, followed by NCP and the chosen partial coordinates. Rows select the smallest centering weight, the distinct weight closest to 0.5, and the largest remaining weight. Plotting never refits a model.

BRM totals ​

These are saved post-hoc position-loss draws. The rows select subjects 720 and 710's load slopes and subject 701's log residual SD. The partial column was checked against the actual stored sampler coordinates; maximum error is 7.2e-15. In particular, scaling the log residual-SD total by its group SD creates the strong curve in the bottom-middle panel.

brms S2Z ​

These are saved S2Z-auto WHMC draws. Rows select subject 720's slope, subject 706's slope and subject 703's log residual-SD contrast. The J subject-labelled contrast values represent J−1 independent directions. The auto column includes the required weighted shift before scaling; it reproduces the actual Qz coordinates to 1.4e-12. S2Z panels show contrasts, while the figure above shows totals.

Recovery sensitivity, validation and limits ​

The gain is not solely a conditional-recovery effect. Excluding all three recovered population coefficients leaves 63 invariant quantities. Online position still gives 238× sampling efficiency and 263× total efficiency against that baseline set. Its total-efficiency advantage over S2Z-auto WHMC remains 2.36×. Ten fresh recovery seeds leave the minimum ESS of both online arms unchanged. Some other rows' minima change when a recovered population coefficient becomes the limiter; the linked sensitivity file records this.

All arms have one chain and 2,000 retained draws, seed 1, with zero sampling divergences. Native NCP hits maximum tree depth on 534 retained transitions. Total NCP has maximum within-chain split R-hat 1.034; the two online arms are about 1.003. Posterior-mean differences from the online-position fit are below three combined estimated MCSEs except one S2Z-NCP WHMC quantity at 3.27. This comparison is a consistency check across 66 quantities, not a proof of convergence.

Native Stan uses 1,000 warmup iterations, adapt_delta 0.8 and maximum tree depth 10. WHMC uses adaptive warmup with Pathfinder initialization. All arms start at the same pooled physical coefficients, with group SDs estimated from subject regressions and both mixture precisions equal to one. Different warmup policies remain part of the sampler comparison. Gradient cost is a computational proxy, not an assertion that every target's gradient has identical wall cost.

The BRM target passed 29 density, gradient, coordinate and recovery checks against an independent Gaussian-integration reference. Ordinary brms and all three S2Z settings also passed density/gradient audits. BRM omits three constant half-Student-t normalizers; adding 3*log(2) aligns its density with the normalized reference without changing gradients or the posterior. Exporting the automatic BRM target and data preserves its density and gradient exactly.

For native Stan, a C++ helper is included when compiling the generated model. A zero-contribution target call counts reverse-mode gradient evaluations, including initialization and warmup. Every retained iteration's count increment was checked against its leapfrog count plus one; final per-process counts include the auto precursor. WHMC counts calls at its target wrapper. Validation-only evaluations are excluded from all fitting costs.

Inspectable model, harness and results ​

The study uses brms revision 73cf607889879cb2a55f50b88d8141d76ff43279, BRM implementation 54cbe3f (the same tree landed as 5f30e53), the repaired WHMC checkout 7aed40b, CmdStan 2.39.0 and CmdStanR 0.9.0. The data revision is d90fc01e6f6fcdced7ee64c9d2ed607d212ec77c.