CDC ww-inference-model: an executable structural port
This page maps the four coupled components of CDC's ww-inference-model onto the @brm formula surface: infection renewal, subpopulation variation, wastewater measurements, and interval-valued count streams. It is an executable structural port, not a drop-in or numerically equivalent reproduction of the current CDC model. The comparison below is verified through transpilation, stanc, and a finite BridgeStan density/gradient by test/cdc_ww_inference.jl; those gates establish a working model artifact, not posterior parity with CDC.
For a smaller introduction to the renewal and shedding pieces, start with the single-catchment wastewater example. A parallel StanBlocks-native case study is available in the StanBlocks documentation.
The current model on the @brm formula surface
This is a genuinely multi-level, multi-stream renewal model. Its composition is on the formula surface, while the sequential numerical kernels are ordinary @deffuns. The seams are:
Shared reference log-Rᵘ —
log_ru_week ~ 1 + dar(week_grid; p=1)is expanded onto the daily latent axis and closed over inside each per-subpopulation cell.Latent subpopulation hierarchy —
I_mat ~ kernel(t_grid, is_reference, logit_I0, initial_growth) do … endbroadcasts only the renewal process. Each cell draws an AR(1) deviation and returns its infection trajectory; the uncovered reference cell suppresses that deviation.Infection feedback and renewal —
renewal_feedbackperforms the carried-state scan inside each cell.Sparse wastewater measurements —
ww_expected_loggathers arbitrary(time, subpopulation, lab)records from the collected latent matrix. The left-censored log-normal likelihood uses a separate lab hierarchy and one LOD per record, so a latent subpopulation may have no wastewater records and several labs may observe the same catchment/time pair.Jurisdiction aggregation —
wsum(I_mat, w)forms a population-weighted infection trajectory.Count streams —
count_interval_meanmaps each stream to arbitrary latent subpopulation weights, a delay PMF, a population multiplier, and a stationary weekly logit-rate trajectory. Observation rows may be daily or inclusive multi-day intervals; a mean-one simplex weekday effect and negative-binomial likelihood finish the component.
The carried-state scans cannot be written as formula statements or kernel control flow. wsum similarly owns the cross-cell matrix multiplication because * at formula level is the Wilkinson interaction operator. These are the interesting numerical kernels, so their exact checked-in Julia definitions are shown rather than left as opaque calls:
StanBlocks.@deffun begin
# CDC's global process: an AR(1) on first differences followed by cumulative
# summation. `eps` has one fewer element than the returned weekly trajectory.
dar_logru(logru1::real, eps::vector[ni], beta::real, sigma::real,
n_weeks::int)::vector[n_weeks] = begin
x::vector[n_weeks]
diff::vector[ni]
x[1] = logru1
if n_weeks >= 2
diff[1] = sigma * eps[1]
x[2] = x[1] + diff[1]
end
for i in 3:n_weeks
diff[i - 1] = beta * diff[i - 2] + sigma * eps[i - 1]
x[i] = x[i - 1] + diff[i - 1]
end
x
end
# Stationary AR(1) deviation around a supplied mean trajectory. CDC uses the
# stationary initial scale for subpopulation R and IHR deviations.
ar1_dev(eps::vector[nt], phi::real, sigma::real)::vector[nt] = begin
d::vector[nt]
d[1] = sigma * eps[1] / sqrt(1.0 - square(phi))
for t in 2:nt
d[t] = phi * d[t - 1] + sigma * eps[t]
end
d
end
# Renewal with infection feedback + exponential seeding.
renewal_feedback(logRt::vector[nt], g::vector[ng], gamma::real, I0::real,
r::real, n_seed::int)::vector[nt] = begin
I::vector[nt]
for t in 1:nt
if t <= n_seed
I[t] = I0 * exp(r * (t - 1))
else
conv = 0.0
for s in 1:ng
prev = t - s >= 1 ? I[t - s] : 0.0
conv = conv + g[s] * prev
end
Rt = exp(logRt[t]) * exp(-gamma * conv)
I[t] = Rt * conv
end
end
I
end
# CDC parameterizes incidence by the per-capita value on the first observed
# day, then back-calculates the beginning of the unobserved growth period.
renewal_from_first_observed(logRt::vector[nt], g::vector[ng], gamma::real,
i_first_obs::real, growth::real,
uot::int)::vector[nt] =
renewal_feedback(logRt, g, gamma,
exp(log(i_first_obs) - uot * growth), growth, uot)
# Normalized triangular shedding trajectory on the log10 viral-load scale,
# matching CDC's `get_vl_trajectory` recurrence.
viral_shedding_trajectory(t_peak::real, viral_peak::real,
duration_shedding::real, n::int)::vector[n] = begin
s::vector[n]
growth = viral_peak / t_peak
wane = viral_peak / (duration_shedding - t_peak)
for t in 1:n
if t <= t_peak
s[t] = exp(log(10.0) * growth * t)
else
log10_load = viral_peak + wane * t_peak - wane * t
s[t] = exp(log(10.0) * (log10_load < 0.0 ? 0.0 : log10_load))
end
end
s / sum(s)
end
# Shedding-load convolution: C(t)=Σ_{k=1}^{nsh} s(k) I(t-k+1).
shed_convolve(I::vector[nt], s::vector[nsh])::vector[nt] = begin
c::vector[nt]
for t in 1:nt
acc = 0.0
for k in 1:nsh
idx = t - k + 1
acc = acc + (idx >= 1 ? s[k] * I[idx] : 0.0)
end
c[t] = acc
end
c
end
# Infection->outcome delay convolution: L(t)=Σ_{k=1}^{nd} d(k) I(t-k+1) (d[1]=lag 0).
delay_convolve(I::vector[nt], d::vector[nd])::vector[nt] = begin
l::vector[nt]
for t in 1:nt
acc = 0.0
for k in 1:nd
idx = t - k + 1
acc = acc + (idx >= 1 ? d[k] * I[idx] : 0.0)
end
l[t] = acc
end
l
end
# Jurisdiction aggregate: population-weighted sum across subpopulation columns
# of the collected infection matrix. `I_mat` is [nt x K] (one column per subpop),
# `w` the population weights; `I_mat * w` is Stan matrix-vector product -> vector[nt].
# (Used by the @brm surface form below, where `*` at the formula level is the
# Wilkinson interaction operator, not matmul — so the matmul lives here.)
wsum(I_mat::matrix[m, K], w::vector[K])::vector[m] = I_mat * w
# Expand one weekly latent value onto a daily row axis through a data index.
weekly_expand(x::vector[nw], week_idx::int[nt])::vector[nt] = x[week_idx]
weekly_expand_columns(x::matrix[nw, S], week_idx::int[nt])::matrix[nt, S] = begin
out::matrix[nt, S]
for t in 1:nt
for stream in 1:S
out[t, stream] = x[week_idx[t], stream]
end
end
out
end
# Build sparse wastewater-record means from the latent subpopulation matrix.
# Records own independent time, subpopulation and lab-site mappings, so a
# latent reference/uncovered population need not have a wastewater series.
ww_expected_log(I_mat::matrix[nt, K], sh::vector[nsh],
sample_time::int[n], sample_subpop::int[n],
sample_lab::int[n], log_lab_mod::vector[nlab],
log10_g::real, mwpd::real)::vector[n] = begin
out::vector[n]
for i in 1:n
shed = 0.0
for lag in 1:nsh
t = sample_time[i] - lag + 1
shed = shed + (t >= 1 ? sh[lag] * I_mat[t, sample_subpop[i]] : 0.0)
end
out[i] = log(10.0) * log10_g + log(shed + 1e-8) - log(mwpd) +
log_lab_mod[sample_lab[i]]
end
out
end
gather_exp(x::vector[nlab], idx::int[n])::vector[n] = exp(x[idx])
take_window(x::vector[n], start_idx::int, width::int)::vector[width] =
x[start_idx:(start_idx + width - 1)]
scale_simplex(x::vector[K], scale::real)::vector[K] = scale * x
hospital_daily_mean(lat::vector[nt], logit_p::vector[nt], dow::int[nt],
dow_effect::vector[7], npop::real)::vector[nt] =
npop * (dow_effect[dow] .* inv_logit(logit_p) .* lat)
# Generic count observation component. Each record names a stream and an
# inclusive time interval. Each stream supplies its own subpopulation
# weights, delay PMF, population multiplier, and weekly rate trajectory.
count_interval_mean(I_mat::matrix[nt, K], subpop_weights::matrix[K, S],
delay::matrix[nd, S], logit_rate::matrix[nt, S],
interval_start::int[n], interval_stop::int[n],
stream_idx::int[n], dow::int[nt],
dow_effect::vector[7], population::vector[S])::vector[n] = begin
out::vector[n]
for record in 1:n
stream = stream_idx[record]
expected = 0.0
for outcome_time in interval_start[record]:interval_stop[record]
delayed = 0.0
for lag in 1:nd
infection_time = outcome_time - lag + 1
if infection_time >= 1
stream_incidence = 0.0
for subpop in 1:K
stream_incidence = stream_incidence +
I_mat[infection_time, subpop] *
subpop_weights[subpop, stream]
end
delayed = delayed + delay[lag, stream] *
inv_logit(logit_rate[infection_time, stream]) *
stream_incidence
end
end
expected = expected + population[stream] *
dow_effect[dow[outcome_time]] * delayed
end
out[record] = expected + 1e-8
end
out
end
# Expected admissions from the aggregate latent infections: day-of-week multiplier
# exp(log_dow), AR(1)-logit IHR inv_logit(logit_p), and jurisdiction population.
ihr_scale(lat::vector[nt], logit_p::vector[nt], log_dow::vector[nt],
npop::real)::vector[nt] = npop * (exp(log_dow) .* inv_logit(logit_p) .* lat)
endThe setup below evaluates those same definitions and the fixture before the four-pane comparison is built.
The declaration is extracted verbatim from cdc_ww_brm_model; the StanBlocks and Stan panes are generated from the resulting BRMI during the docs build.
CDC ww-inference structural port (@brm formula surface)function cdc_ww_brm_model(df = cdc_ww_brm_fixture())
@brm df begin
# Shared epidemic parameters. `dar` owns the zero-start differenced-AR
# innovations while the intercept remains the initial weekly log-Rᵘ level.
gamma ~ LogNormal(-4.0, 0.5)
phi_delta ~ Beta(2.0, 8.0)
sigma_delta ~ Exponential(4.0)
log10_g ~ Normal(12.0, 1.0)
t_peak ~ LogNormal(log(4.0), 0.25)
viral_peak ~ Normal(6.0, 1.0)
shed_tail ~ LogNormal(log(14.0), 0.3)
dur_shed = t_peak + shed_tail
shedding = viral_shedding_trajectory(t_peak, viral_peak, dur_shed, nsh)
log_ru_week ~ 1 + dar(week_grid; p = 1)
effect(log_ru_week, Intercept) ~ Normal(0.0, 0.5)
log_ru = weekly_expand(log_ru_week, week_idx)
# Hierarchical incidence and initial growth are defined over the latent
# subpopulation axis, including the uncovered reference population.
logit_I0 ~ 1 + (1 | initial | subpopulation)
initial_growth ~ 1 + (1 | growth | subpopulation)
effect(logit_I0, Intercept) ~ Normal(-9.0, 1.0)
effect(initial_growth, Intercept) ~ Normal(0.0, 0.05)
sd(:, initial) ~ Normal(0.0, 0.75)
sd(:, growth) ~ Normal(0.0, 0.05)
# The kernel produces latent trajectories only. Observation frames are
# mapped onto this matrix after the cell has been collected.
I_mat ~ kernel(t_grid, is_reference, logit_I0, initial_growth) do ts, isref, lI0, growth_i
eps_d::vector[n_total] ~ std_normal()
delta_k = ar1_dev(eps_d, phi_delta, sigma_delta)
logRt_k = log_ru + (1.0 - isref) * delta_k
renewal_from_first_observed(logRt_k, g, gamma,
inv_logit(lI0), growth_i, uot)
end
# Sparse wastewater records have independent time, subpopulation, lab
# and record-specific LOD mappings. Multiple labs may observe the same
# catchment/time pair; the reference population may have no records.
log_lab_mod ~ 0 + (1 | ww_scale | ww_lab)
log_sigma_ww ~ 1 + (1 | ww_noise | ww_lab)
effect(log_sigma_ww, Intercept) ~ Normal(log(0.35), 0.5)
sd(:, ww_scale) ~ Normal(0.0, 0.5)
sd(:, ww_noise) ~ Normal(0.0, 0.35)
ww_mu = ww_expected_log(I_mat, shedding, ww_time, ww_subpop, ww_lab_idx,
log_lab_mod, log10_g, mwpd)
ww_sigma = gather_exp(log_sigma_ww, ww_lab_idx)
ww_log_conc ~ censored(Normal(ww_mu, ww_sigma); lower = ww_lod)
# jurisdiction aggregate infections (population-weighted across subpops)
I_agg = wsum(I_mat, w)
# Count-stream module. Each stream gets a hierarchical rate center and a
# stationary weekly AR deviation. Observation rows can be daily or span
# arbitrary inclusive intervals and can map different latent catchments.
phi_count ~ Beta(2.0, 8.0)
sigma_count ~ Exponential(20.0)
count_rate_center ~ 1 + (1 | count_rate | count_stream)
effect(count_rate_center, Intercept) ~ Normal(-4.6, 0.3)
sd(:, count_rate) ~ Normal(0.0, 0.5)
count_rate_week ~ kernel(count_week_grid, count_rate_center) do weeks, center
eps_count::vector[n_weeks] ~ std_normal()
center + ar1_dev(eps_count, phi_count, sigma_count)
end
count_rate_daily = weekly_expand_columns(count_rate_week, week_idx)
dow_share ~ Dirichlet(7, 5.0)
dow_effect = scale_simplex(dow_share, 7.0)
log_phi_count ~ 1 + (1 | count_dispersion | count_stream)
effect(log_phi_count, Intercept) ~ Normal(log(15.0), 0.5)
sd(:, count_dispersion) ~ Normal(0.0, 0.5)
count_mu = count_interval_mean(I_mat, count_subpop_weights, count_delay,
count_rate_daily, count_start, count_stop,
count_stream_idx, dow, dow_effect,
count_population)
count_phi = gather_exp(log_phi_count, count_stream_idx)
count ~ NegativeBinomial2(count_mu, count_phi)
# Forecast contract: named deterministic carriers use the same fitted
# latent state and observation mappings but are not calibration outcomes.
forecast_infections = take_window(I_agg, uot + nt + 1, ht)
forecast_count_mean = count_interval_mean(
I_mat, count_subpop_weights, count_delay, count_rate_daily,
forecast_count_start, forecast_count_stop,
forecast_count_stream_idx, dow, dow_effect, count_population)
forecast_ww_log_mean = ww_expected_log(
I_mat, shedding, forecast_ww_time, forecast_ww_subpop,
forecast_ww_lab_idx, log_lab_mod, log10_g, mwpd)
end
endBRMI:
gamma ~ LogNormal(-4.0, 0.5)
phi_delta ~ Beta(2.0, 8.0)
sigma_delta ~ Exponential(4.0)
log10_g ~ Normal(12.0, 1.0)
t_peak ~ LogNormal(log(4.0), 0.25)
viral_peak ~ Normal(6.0, 1.0)
shed_tail ~ LogNormal(log(14.0), 0.3)
:dur_shed = t_peak + shed_tail
nsh: data (eltype=Int64, n=1)
:shedding = viral_shedding_trajectory(t_peak, viral_peak, dur_shed, nsh)
week_grid: data (eltype=Float64, n=22)
log_ru_week ~ 1 + dar(week_grid; p=1)
effect(log_ru_week, Intercept) ~ Normal(0.0, 0.5)
week_idx: data (eltype=Int64, n=148)
:log_ru = weekly_expand(log_ru_week, week_idx)
subpopulation: data (eltype=String, n=3)
logit_I0 ~ 1 + (1 | initial | subpopulation)
initial_growth ~ 1 + (1 | growth | subpopulation)
effect(logit_I0, Intercept) ~ Normal(-9.0, 1.0)
effect(initial_growth, Intercept) ~ Normal(0.0, 0.05)
effect(sd, initial) ~ Normal(0.0, 0.75)
effect(sd, growth) ~ Normal(0.0, 0.05)
t_grid: data (eltype=Vector{Float64}, n=3)
is_reference: data (eltype=Float64, n=3)
I_mat ~ kernel((ts, isref, lI0, growth_i)->begin
#= brm-docs-example.jl:31 =#
eps_d::vector[n_total] ~ std_normal()
#= brm-docs-example.jl:32 =#
delta_k = ar1_dev(eps_d, phi_delta, sigma_delta)
#= brm-docs-example.jl:33 =#
logRt_k = log_ru + (1.0 - isref) * delta_k
#= brm-docs-example.jl:34 =#
renewal_from_first_observed(logRt_k, g, gamma, inv_logit(lI0), growth_i, uot)
end, t_grid, is_reference, logit_I0, initial_growth)
ww_lab: data (eltype=String, n=3)
log_lab_mod ~ 0 + (1 | ww_scale | ww_lab)
log_sigma_ww ~ 1 + (1 | ww_noise | ww_lab)
effect(log_sigma_ww, Intercept) ~ Normal(log(0.35), 0.5)
effect(sd, ww_scale) ~ Normal(0.0, 0.5)
effect(sd, ww_noise) ~ Normal(0.0, 0.35)
ww_time: data (eltype=Int64, n=36)
ww_subpop: data (eltype=Int64, n=36)
ww_lab_idx: data (eltype=Int64, n=36)
mwpd: data (eltype=Float64, n=1)
:ww_mu = ww_expected_log(I_mat, shedding, ww_time, ww_subpop, ww_lab_idx, log_lab_mod, log10_g, mwpd)
:ww_sigma = gather_exp(log_sigma_ww, ww_lab_idx)
ww_lod: data (eltype=Float64, n=36)
ww_log_conc ~ censored(Normal(ww_mu, ww_sigma); lower=ww_lod)
w: data (eltype=Float64, n=3)
:I_agg = wsum(I_mat, w)
phi_count ~ Beta(2.0, 8.0)
sigma_count ~ Exponential(20.0)
count_stream: data (eltype=String, n=2)
count_rate_center ~ 1 + (1 | count_rate | count_stream)
effect(count_rate_center, Intercept) ~ Normal(-4.6, 0.3)
effect(sd, count_rate) ~ Normal(0.0, 0.5)
count_week_grid: data (eltype=Vector{Float64}, n=2)
count_rate_week ~ kernel((weeks, center)->begin
#= brm-docs-example.jl:63 =#
eps_count::vector[n_weeks] ~ std_normal()
#= brm-docs-example.jl:64 =#
center + ar1_dev(eps_count, phi_count, sigma_count)
end, count_week_grid, count_rate_center)
:count_rate_daily = weekly_expand_columns(count_rate_week, week_idx)
dow_share ~ Dirichlet(7, 5.0)
:dow_effect = scale_simplex(dow_share, 7.0)
log_phi_count ~ 1 + (1 | count_dispersion | count_stream)
effect(log_phi_count, Intercept) ~ Normal(log(15.0), 0.5)
effect(sd, count_dispersion) ~ Normal(0.0, 0.5)
count_subpop_weights: data (eltype=Float64, n=6)
count_delay: data (eltype=Float64, n=28)
count_start: data (eltype=Int64, n=96)
count_stop: data (eltype=Int64, n=96)
count_stream_idx: data (eltype=Int64, n=96)
dow: data (eltype=Int64, n=148)
count_population: data (eltype=Float64, n=2)
:count_mu = count_interval_mean(I_mat, count_subpop_weights, count_delay, count_rate_daily, count_start, count_stop, count_stream_idx, dow, dow_effect, count_population)
:count_phi = gather_exp(log_phi_count, count_stream_idx)
count ~ NegativeBinomial2(count_mu, count_phi)
uot: data (eltype=Int64, n=1)
nt: data (eltype=Int64, n=1)
ht: data (eltype=Int64, n=1)
:forecast_infections = take_window(I_agg, (uot + nt + 1), ht)
forecast_count_start: data (eltype=Int64, n=16)
forecast_count_stop: data (eltype=Int64, n=16)
forecast_count_stream_idx: data (eltype=Int64, n=16)
:forecast_count_mean = count_interval_mean(I_mat, count_subpop_weights, count_delay, count_rate_daily, forecast_count_start, forecast_count_stop, forecast_count_stream_idx, dow, dow_effect, count_population)
forecast_ww_time: data (eltype=Int64, n=6)
forecast_ww_subpop: data (eltype=Int64, n=6)
forecast_ww_lab_idx: data (eltype=Int64, n=6)
:forecast_ww_log_mean = ww_expected_log(I_mat, shedding, forecast_ww_time, forecast_ww_subpop, forecast_ww_lab_idx, log_lab_mod, log10_g, mwpd)
g: data (eltype=Float64, n=12)
n_total: data (eltype=Int64, n=1)
n_weeks: data (eltype=Int64, n=1)SBBRMI with data keys = [:count, :count_delay, :count_population, :count_start, :count_stop, :count_stream_idx, :count_subpop_weights, :count_week_grid, :dow, :dow_share_alpha, :forecast_count_start, :forecast_count_stop, :forecast_count_stream_idx, :forecast_ww_lab_idx, :forecast_ww_subpop, :forecast_ww_time, :g, :ht, :is_reference, :kernel_nsub_I_mat, :kernel_nsub_count_rate_week, :mwpd, :n_count_stream, :n_subpopulation, :n_terms_count_dispersion_count_stream, :n_terms_count_rate_count_stream, :n_terms_growth_subpopulation, :n_terms_initial_subpopulation, :n_terms_ww_noise_ww_lab, :n_terms_ww_scale_ww_lab, :n_total, :n_weeks, :n_ww_lab, :nsh, :nt, :subpopulation_idx, :t_grid, :uot, :w, :week_grid, :week_idx, :ww_lab_idx, :ww_lod, :ww_log_conc, :ww_subpop, :ww_time]
configured submodels:
ranef_correlated_draws_generic_configured_1 = Base.merge(BayesianRegressionModels.ranef_correlated_draws_generic, quote
tau ~ normal(0.0, 0.75; n = n_terms, lower = 0.0)
end)
ranef_correlated_draws_generic_configured_2 = Base.merge(BayesianRegressionModels.ranef_correlated_draws_generic, quote
tau ~ normal(0.0, 0.05; n = n_terms, lower = 0.0)
end)
ranef_correlated_draws_generic_configured_3 = Base.merge(BayesianRegressionModels.ranef_correlated_draws_generic, quote
tau ~ normal(0.0, 0.5; n = n_terms, lower = 0.0)
end)
ranef_correlated_draws_generic_configured_4 = Base.merge(BayesianRegressionModels.ranef_correlated_draws_generic, quote
tau ~ normal(0.0, 0.35; n = n_terms, lower = 0.0)
end)
emitted @slic body:
begin
b_initial_subpopulation ~ ranef_correlated_draws_generic_configured_1(; group_idx = subpopulation_idx, n_groups = n_subpopulation, n_terms = n_terms_initial_subpopulation, lkj_eta = 1.0)
b_growth_subpopulation ~ ranef_correlated_draws_generic_configured_2(; group_idx = subpopulation_idx, n_groups = n_subpopulation, n_terms = n_terms_growth_subpopulation, lkj_eta = 1.0)
b_ww_scale_ww_lab ~ ranef_correlated_draws_generic_configured_3(; group_idx = ww_lab_idx, n_groups = n_ww_lab, n_terms = n_terms_ww_scale_ww_lab, lkj_eta = 1.0)
b_ww_noise_ww_lab ~ ranef_correlated_draws_generic_configured_4(; group_idx = ww_lab_idx, n_groups = n_ww_lab, n_terms = n_terms_ww_noise_ww_lab, lkj_eta = 1.0)
b_count_rate_count_stream ~ ranef_correlated_draws_generic_configured_3(; group_idx = count_stream_idx, n_groups = n_count_stream, n_terms = n_terms_count_rate_count_stream, lkj_eta = 1.0)
b_count_dispersion_count_stream ~ ranef_correlated_draws_generic_configured_3(; group_idx = count_stream_idx, n_groups = n_count_stream, n_terms = n_terms_count_dispersion_count_stream, lkj_eta = 1.0)
gamma ~ lognormal(-4.0, 0.5)
phi_delta ~ beta(2.0, 8.0)
sigma_delta ~ exponential(1.0 ./ 4.0)
log10_g ~ normal(12.0, 1.0)
t_peak ~ lognormal(1.3862943611198906, 0.25)
viral_peak ~ normal(6.0, 1.0)
shed_tail ~ lognormal(2.6390573296152584, 0.3)
dur_shed = (+)(t_peak, shed_tail)
shedding = (Main.cdc_ww.viral_shedding_trajectory)(t_peak, viral_peak, dur_shed, nsh)
X_log_ru_week = hcat(rep_vector(1.0, num_elements(week_grid)))
pop_log_ru_week ~ _popefs_normal(; X = X_log_ru_week, beta_loc = [0.0], beta_scale = [0.5])
dar_log_ru_week_week_grid ~ _sb_dar1(; time = week_grid)
log_ru_week = pop_log_ru_week + dar_log_ru_week_week_grid
log_ru = (Main.cdc_ww.weekly_expand)(log_ru_week, week_idx)
X_logit_I0 = hcat(rep_vector(1.0, num_elements(subpopulation_idx)))
pop_logit_I0 ~ _popefs_normal(; X = X_logit_I0, beta_loc = [-9.0], beta_scale = [1.0])
r_logit_I0_initial_subpopulation = b_initial_subpopulation[subpopulation_idx, 1]
logit_I0 = pop_logit_I0 + r_logit_I0_initial_subpopulation
X_initial_growth = hcat(rep_vector(1.0, num_elements(subpopulation_idx)))
pop_initial_growth ~ _popefs_normal(; X = X_initial_growth, beta_loc = [0.0], beta_scale = [0.05])
r_initial_growth_growth_subpopulation = b_growth_subpopulation[subpopulation_idx, 1]
initial_growth = pop_initial_growth + r_initial_growth_growth_subpopulation
I_mat ~ plate(t_grid, is_reference, logit_I0, initial_growth; outer = (kernel_nsub_I_mat,)) do ts, isref, lI0, growth_i
#= brm-docs-example.jl:31 =#
eps_d::vector[n_total] ~ std_normal()
#= brm-docs-example.jl:32 =#
delta_k = ar1_dev(eps_d, phi_delta, sigma_delta)
#= brm-docs-example.jl:33 =#
logRt_k = log_ru + (1.0 - isref) * delta_k
#= brm-docs-example.jl:34 =#
renewal_from_first_observed(logRt_k, g, gamma, inv_logit(lI0), growth_i, uot)
end
r_log_lab_mod_ww_scale_ww_lab = b_ww_scale_ww_lab[ww_lab_idx, 1]
log_lab_mod = r_log_lab_mod_ww_scale_ww_lab
X_log_sigma_ww = hcat(rep_vector(1.0, num_elements(ww_lab_idx)))
pop_log_sigma_ww ~ _popefs_normal(; X = X_log_sigma_ww, beta_loc = [-1.0498221244986778], beta_scale = [0.5])
r_log_sigma_ww_ww_noise_ww_lab = b_ww_noise_ww_lab[ww_lab_idx, 1]
log_sigma_ww = pop_log_sigma_ww + r_log_sigma_ww_ww_noise_ww_lab
ww_mu = (Main.cdc_ww.ww_expected_log)(I_mat, shedding, ww_time, ww_subpop, ww_lab_idx, log_lab_mod, log10_g, mwpd)
ww_sigma = (Main.cdc_ww.gather_exp)(log_sigma_ww, ww_lab_idx)
ww_log_conc ~ censored(normal, ww_mu, ww_sigma; lower = ww_lod)
I_agg = (Main.cdc_ww.wsum)(I_mat, w)
phi_count ~ beta(2.0, 8.0)
sigma_count ~ exponential(1.0 ./ 20.0)
X_count_rate_center = hcat(rep_vector(1.0, num_elements(count_stream_idx)))
pop_count_rate_center ~ _popefs_normal(; X = X_count_rate_center, beta_loc = [-4.6], beta_scale = [0.3])
r_count_rate_center_count_rate_count_stream = b_count_rate_count_stream[count_stream_idx, 1]
count_rate_center = pop_count_rate_center + r_count_rate_center_count_rate_count_stream
count_rate_week ~ plate(count_week_grid, count_rate_center; outer = (kernel_nsub_count_rate_week,)) do weeks, center
#= brm-docs-example.jl:63 =#
eps_count::vector[n_weeks] ~ std_normal()
#= brm-docs-example.jl:64 =#
center + ar1_dev(eps_count, phi_count, sigma_count)
end
count_rate_daily = (Main.cdc_ww.weekly_expand_columns)(count_rate_week, week_idx)
dow_share ~ dirichlet(dow_share_alpha)
dow_effect = (Main.cdc_ww.scale_simplex)(dow_share, 7.0)
X_log_phi_count = hcat(rep_vector(1.0, num_elements(count_stream_idx)))
pop_log_phi_count ~ _popefs_normal(; X = X_log_phi_count, beta_loc = [2.70805020110221], beta_scale = [0.5])
r_log_phi_count_count_dispersion_count_stream = b_count_dispersion_count_stream[count_stream_idx, 1]
log_phi_count = pop_log_phi_count + r_log_phi_count_count_dispersion_count_stream
count_mu = (Main.cdc_ww.count_interval_mean)(I_mat, count_subpop_weights, count_delay, count_rate_daily, count_start, count_stop, count_stream_idx, dow, dow_effect, count_population)
count_phi = (Main.cdc_ww.gather_exp)(log_phi_count, count_stream_idx)
count ~ neg_binomial_2(count_mu, count_phi)
forecast_infections = (Main.cdc_ww.take_window)(I_agg, (+)(uot, nt, 1), ht)
forecast_count_mean = (Main.cdc_ww.count_interval_mean)(I_mat, count_subpop_weights, count_delay, count_rate_daily, forecast_count_start, forecast_count_stop, forecast_count_stream_idx, dow, dow_effect, count_population)
forecast_ww_log_mean = (Main.cdc_ww.ww_expected_log)(I_mat, shedding, forecast_ww_time, forecast_ww_subpop, forecast_ww_lab_idx, log_lab_mod, log10_g, mwpd)
endfunctions {
vector viral_shedding_trajectory(
real t_peak,
real viral_peak,
real duration_shedding,
int n
) {
vector[n] s;
real growth = (viral_peak / t_peak);
real wane = (viral_peak / (duration_shedding - t_peak));
for(t in 1:n) {
if((t <= t_peak)) {
s[t] = exp((log(10.0) * growth * t));
} else {
real log10_load = ((viral_peak + (wane * t_peak)) - (wane * t));
s[t] = exp((log(10.0) * ((log10_load < 0.0) ? 0.0 : log10_load)));
}
}
return (s / sum(s));
}
matrix hcat(vector x) {
int n = dims(x)[1];
return to_matrix(x, n, 1);
}
vector differenced_ar1_path(
real beta,
real sigma,
vector z
) {
int n = dims(z)[1];
vector[(n + 1)] x = rep_vector(0.0, (n + 1));
real increment = 0.0;
if((n > 0)) {
for(t in 1:n) {
increment = ((beta * increment) + (sigma * z[t]));
x[(t + 1)] = (x[t] + increment);
}
}
return x;
}
vector weekly_expand(
vector x,
array[] int week_idx
) {
int nt = dims(week_idx)[1];
return x[week_idx];
}
vector std_normal_vector_rng(
int anontok__1
) {
int n = anontok__1;
if((n == 0)) {
vector[n] rv;
return rv;
} else {
return to_vector(normal_rng(rep_vector(0, n), 1));
}
}
vector ar1_dev(
vector eps,
real phi,
real sigma
) {
int nt = dims(eps)[1];
vector[nt] d;
d[1] = ((sigma * eps[1]) / sqrt((1.0 - square(phi))));
for(t in 2:nt) {
d[t] = ((phi * d[(t - 1)]) + (sigma * eps[t]));
}
return d;
}
vector renewal_from_first_observed(
vector logRt,
vector g,
real gamma,
real i_first_obs,
real growth,
int uot
) {
int nt = dims(logRt)[1];
return renewal_feedback(logRt, g, gamma, exp((log(i_first_obs) - (uot * growth))), growth, uot);
}
vector renewal_feedback(
vector logRt,
vector g,
real gamma,
real I0,
real r,
int n_seed
) {
int nt = dims(logRt)[1];
int ng = dims(g)[1];
vector[nt] I;
for(t in 1:nt) {
if((t <= n_seed)) {
I[t] = (I0 * exp((r * (t - 1))));
} else {
real conv = 0.0;
for(s in 1:ng) {
real prev = (((t - s) >= 1) ? I[(t - s)] : 0.0);
conv = (conv + (g[s] * prev));
}
real Rt = (exp(logRt[t]) * exp(((-gamma) * conv)));
I[t] = (Rt * conv);
}
}
return I;
}
vector ww_expected_log(
matrix I_mat,
vector sh,
array[] int sample_time,
array[] int sample_subpop,
array[] int sample_lab,
vector log_lab_mod,
real log10_g,
real mwpd
) {
int nsh = dims(sh)[1];
int n = dims(sample_time)[1];
if (dims(sample_subpop)[1] != n) reject("ww_expected_log: dim mismatch — `sample_subpop` dim 1 (= ", dims(sample_subpop)[1], ") does not match `n` (= ", n, "), inferred from `sample_time` dim 1. `n` sizes: `sample_time` dim 1 (= ", dims(sample_time)[1], "), `sample_subpop` dim 1 (= ", dims(sample_subpop)[1], "), `sample_lab` dim 1 (= ", dims(sample_lab)[1], ").");
if (dims(sample_lab)[1] != n) reject("ww_expected_log: dim mismatch — `sample_lab` dim 1 (= ", dims(sample_lab)[1], ") does not match `n` (= ", n, "), inferred from `sample_time` dim 1. `n` sizes: `sample_time` dim 1 (= ", dims(sample_time)[1], "), `sample_subpop` dim 1 (= ", dims(sample_subpop)[1], "), `sample_lab` dim 1 (= ", dims(sample_lab)[1], ").");
vector[n] out;
for(i in 1:n) {
real shed = 0.0;
for(lag in 1:nsh) {
int t = ((sample_time[i] - lag) + 1);
shed = (shed + ((t >= 1) ? (sh[lag] * I_mat[t, sample_subpop[i]]) : 0.0));
}
out[i] = ((((log(10.0) * log10_g) + log((shed + 1.0e-8))) - log(mwpd)) + log_lab_mod[sample_lab[i]]);
}
return out;
}
vector gather_exp(vector x, array[] int idx) {
int n = dims(idx)[1];
return exp(x[idx]);
}
real lower_clamping_normal_lpdf(
vector y,
vector lo,
vector args1,
vector args2
) {
return sum(lower_clamping_normal_lpdfs(y, lo, args1, args2));
}
vector lower_clamping_normal_lpdfs(
vector y,
vector lo,
vector args1,
vector args2
) {
int n = dims(y)[1];
return jbroadcasted_lower_clamping_lpdf_normal(y, lo, args1, args2);
}
vector jbroadcasted_lower_clamping_lpdf_normal(
vector x1,
vector x3,
vector x4,
vector x5
) {
int n = dims(x1)[1];
vector[n] rv;
for(i in 1:n) {
rv[i] = lower_clamping_normal_lpdf(broadcasted_getindex(x1, i) |
broadcasted_getindex(x3, i),
broadcasted_getindex(x4, i),
broadcasted_getindex(x5, i)
);
}
return rv;
}
real lower_clamping_normal_lpdf(
real y,
real lo,
real args1,
real args2
) {
array[1] real rv;
rv[1] = negative_infinity();
if((y == lo)) {
rv[1] = normal_lcdf_stable(lo, args1, args2);
} else {
if((y > lo)) {
rv[1] = normal_lpdf(y | args1, args2);
}
}
return rv[1];
}
real normal_lcdf_stable(
real x,
real loc,
real scale
) {
return (log(erfc(((-(x - loc)) / (scale * sqrt(2.0))))) - log(2.0));
}
real broadcasted_getindex(vector x, int i) {
return x[i];
}
vector lower_clamping_vector_normal_rng(
int anontok__1,
vector lo,
vector args1,
vector args2
) {
int n = anontok__1;
return jbroadcasted_lower_clamping_cell_rng_normal_rng(rep_vector(0.0, n), lo, args1, args2);
}
vector jbroadcasted_lower_clamping_cell_rng_normal_rng(
vector x1,
vector x3,
vector x4,
vector x5
) {
int n = dims(x1)[1];
vector[n] rv;
for(i in 1:n) {
rv[i] = lower_clamping_cell_normal_rng(
broadcasted_getindex(x1, i),
broadcasted_getindex(x3, i),
broadcasted_getindex(x4, i),
broadcasted_getindex(x5, i)
);
}
return rv;
}
real lower_clamping_cell_normal_rng(
real dummy,
real lo,
real args1,
real args2
) {
return lower_clamping_normal_rng(lo, args1, args2);
}
real lower_clamping_normal_rng(
real lo,
real args1,
real args2
) {
vector[1] draw;
draw[1] = normal_rng(args1, args2);
if((draw[1] < lo)) {
draw[1] = lo;
}
return draw[1];
}
vector wsum(
matrix I_mat,
vector w
) {
int m = dims(I_mat)[1];
int K = dims(I_mat)[2];
if (dims(w)[1] != K) reject("wsum: dim mismatch — `w` dim 1 (= ", dims(w)[1], ") does not match `K` (= ", K, "), inferred from `I_mat` dim 2. `K` sizes: `I_mat` dim 2 (= ", dims(I_mat)[2], "), `w` dim 1 (= ", dims(w)[1], ").");
return (I_mat * w);
}
matrix weekly_expand_columns(
matrix x,
array[] int week_idx
) {
int S = dims(x)[2];
int nt = dims(week_idx)[1];
matrix[nt, S] out;
for(t in 1:nt) {
for(stream in 1:S) {
out[t, stream] = x[week_idx[t], stream];
}
}
return out;
}
vector scale_simplex(vector x, real scale) {
int K = dims(x)[1];
return (scale * x);
}
vector count_interval_mean(
matrix I_mat,
matrix subpop_weights,
matrix delay,
matrix logit_rate,
array[] int interval_start,
array[] int interval_stop,
array[] int stream_idx,
array[] int dow,
vector dow_effect,
vector population
) {
int nt = dims(I_mat)[1];
int K = dims(I_mat)[2];
int S = dims(subpop_weights)[2];
int nd = dims(delay)[1];
int n = dims(interval_start)[1];
if (dims(logit_rate)[1] != nt) reject("count_interval_mean: dim mismatch — `logit_rate` dim 1 (= ", dims(logit_rate)[1], ") does not match `nt` (= ", nt, "), inferred from `I_mat` dim 1. `nt` sizes: `I_mat` dim 1 (= ", dims(I_mat)[1], "), `logit_rate` dim 1 (= ", dims(logit_rate)[1], "), `dow` dim 1 (= ", dims(dow)[1], ").");
if (dims(dow)[1] != nt) reject("count_interval_mean: dim mismatch — `dow` dim 1 (= ", dims(dow)[1], ") does not match `nt` (= ", nt, "), inferred from `I_mat` dim 1. `nt` sizes: `I_mat` dim 1 (= ", dims(I_mat)[1], "), `logit_rate` dim 1 (= ", dims(logit_rate)[1], "), `dow` dim 1 (= ", dims(dow)[1], ").");
if (dims(subpop_weights)[1] != K) reject("count_interval_mean: dim mismatch — `subpop_weights` dim 1 (= ", dims(subpop_weights)[1], ") does not match `K` (= ", K, "), inferred from `I_mat` dim 2. `K` sizes: `I_mat` dim 2 (= ", dims(I_mat)[2], "), `subpop_weights` dim 1 (= ", dims(subpop_weights)[1], ").");
if (dims(delay)[2] != S) reject("count_interval_mean: dim mismatch — `delay` dim 2 (= ", dims(delay)[2], ") does not match `S` (= ", S, "), inferred from `subpop_weights` dim 2. `S` sizes: `subpop_weights` dim 2 (= ", dims(subpop_weights)[2], "), `delay` dim 2 (= ", dims(delay)[2], "), `logit_rate` dim 2 (= ", dims(logit_rate)[2], "), `population` dim 1 (= ", dims(population)[1], ").");
if (dims(logit_rate)[2] != S) reject("count_interval_mean: dim mismatch — `logit_rate` dim 2 (= ", dims(logit_rate)[2], ") does not match `S` (= ", S, "), inferred from `subpop_weights` dim 2. `S` sizes: `subpop_weights` dim 2 (= ", dims(subpop_weights)[2], "), `delay` dim 2 (= ", dims(delay)[2], "), `logit_rate` dim 2 (= ", dims(logit_rate)[2], "), `population` dim 1 (= ", dims(population)[1], ").");
if (dims(population)[1] != S) reject("count_interval_mean: dim mismatch — `population` dim 1 (= ", dims(population)[1], ") does not match `S` (= ", S, "), inferred from `subpop_weights` dim 2. `S` sizes: `subpop_weights` dim 2 (= ", dims(subpop_weights)[2], "), `delay` dim 2 (= ", dims(delay)[2], "), `logit_rate` dim 2 (= ", dims(logit_rate)[2], "), `population` dim 1 (= ", dims(population)[1], ").");
if (dims(interval_stop)[1] != n) reject("count_interval_mean: dim mismatch — `interval_stop` dim 1 (= ", dims(interval_stop)[1], ") does not match `n` (= ", n, "), inferred from `interval_start` dim 1. `n` sizes: `interval_start` dim 1 (= ", dims(interval_start)[1], "), `interval_stop` dim 1 (= ", dims(interval_stop)[1], "), `stream_idx` dim 1 (= ", dims(stream_idx)[1], ").");
if (dims(stream_idx)[1] != n) reject("count_interval_mean: dim mismatch — `stream_idx` dim 1 (= ", dims(stream_idx)[1], ") does not match `n` (= ", n, "), inferred from `interval_start` dim 1. `n` sizes: `interval_start` dim 1 (= ", dims(interval_start)[1], "), `interval_stop` dim 1 (= ", dims(interval_stop)[1], "), `stream_idx` dim 1 (= ", dims(stream_idx)[1], ").");
vector[n] out;
for(record in 1:n) {
int stream = stream_idx[record];
real expected = 0.0;
for(outcome_time in interval_start[record]:interval_stop[record]) {
real delayed = 0.0;
for(lag in 1:nd) {
int infection_time = ((outcome_time - lag) + 1);
if((infection_time >= 1)) {
real stream_incidence = 0.0;
for(subpop in 1:K) {
stream_incidence = (stream_incidence + (I_mat[infection_time, subpop] * subpop_weights[subpop, stream]));
}
delayed = (delayed + (delay[lag, stream] * inv_logit(logit_rate[infection_time, stream]) * stream_incidence));
}
}
expected = (expected + (population[stream] * dow_effect[dow[outcome_time]] * delayed));
}
out[record] = (expected + 1.0e-8);
}
return out;
}
vector neg_binomial_2_lpmfs(
array[] int obs,
vector mu,
vector phi
) {
return jbroadcasted_neg_binomial_2_lpmfs(obs, mu, phi);
}
vector jbroadcasted_neg_binomial_2_lpmfs(
array[] int x1,
vector x2,
vector x3
) {
int n = dims(x1)[1];
vector[n] rv;
for(i in 1:n) {
rv[i] = neg_binomial_2_lpmfs(
broadcasted_getindex(x1, i),
broadcasted_getindex(x2, i),
broadcasted_getindex(x3, i)
);
}
return rv;
}
real neg_binomial_2_lpmfs(
int args1,
real args2,
real args3
) {
return neg_binomial_2_lpmf(args1 | args2, args3);
}
int broadcasted_getindex(array[] int x, int i) {
return x[i];
}
array[] int neg_binomial_2_int_rng(
int anontok__1,
vector a,
vector b
) {
int n = anontok__1;
if((n == 0)) {
array[n] int rv;
return rv;
} else {
return neg_binomial_2_rng(a, b);
}
}
vector take_window(
vector x,
int start_idx,
int width
) {
return x[start_idx:((start_idx + width) - 1)];
}
}
data {
int n_terms_initial_subpopulation;
int n_subpopulation;
int n_terms_growth_subpopulation;
int n_terms_ww_scale_ww_lab;
int n_ww_lab;
int n_terms_ww_noise_ww_lab;
int n_terms_count_rate_count_stream;
int n_count_stream;
int n_terms_count_dispersion_count_stream;
int nsh;
int week_grid_n;
vector[week_grid_n] week_grid;
int week_idx_n;
array[week_idx_n] int week_idx;
int subpopulation_idx_n;
array[subpopulation_idx_n] int subpopulation_idx;
int n_total;
int kernel_nsub_I_mat;
int is_reference_n;
vector[is_reference_n] is_reference;
int g_n;
vector[g_n] g;
int uot;
int ww_lab_idx_n;
array[ww_lab_idx_n] int ww_lab_idx;
int ww_time_n;
array[ww_time_n] int ww_time;
int ww_subpop_n;
array[ww_subpop_n] int ww_subpop;
real mwpd;
int ww_log_conc_n;
vector[ww_log_conc_n] ww_log_conc;
int ww_lod_n;
vector[ww_lod_n] ww_lod;
int w_n;
vector[w_n] w;
int count_stream_idx_n;
array[count_stream_idx_n] int count_stream_idx;
int n_weeks;
int kernel_nsub_count_rate_week;
int dow_share_alpha_n;
vector[dow_share_alpha_n] dow_share_alpha;
int count_subpop_weights_m;
int count_subpop_weights_n;
matrix[count_subpop_weights_m, count_subpop_weights_n] count_subpop_weights;
int count_delay_m;
int count_delay_n;
matrix[count_delay_m, count_delay_n] count_delay;
int count_start_n;
array[count_start_n] int count_start;
int count_stop_n;
array[count_stop_n] int count_stop;
int dow_n;
array[dow_n] int dow;
int count_population_n;
vector[count_population_n] count_population;
int count_n;
array[count_n] int count;
int ht;
int nt;
int forecast_count_stream_idx_n;
int forecast_count_start_n;
array[forecast_count_start_n] int forecast_count_start;
int forecast_count_stop_n;
array[forecast_count_stop_n] int forecast_count_stop;
array[forecast_count_stream_idx_n] int forecast_count_stream_idx;
int forecast_ww_lab_idx_n;
int forecast_ww_time_n;
array[forecast_ww_time_n] int forecast_ww_time;
int forecast_ww_subpop_n;
array[forecast_ww_subpop_n] int forecast_ww_subpop;
array[forecast_ww_lab_idx_n] int forecast_ww_lab_idx;
}
transformed data {
matrix[num_elements(week_grid), 1] X_log_ru_week = hcat(rep_vector(1.0, num_elements(week_grid)));
int pop_log_ru_week_n_covariates = 1;
int dar_log_ru_week_week_grid_n_innov = (num_elements(week_grid) - 1);
matrix[num_elements(subpopulation_idx), 1] X_logit_I0 = hcat(rep_vector(1.0, num_elements(subpopulation_idx)));
int pop_logit_I0_n_covariates = 1;
matrix[num_elements(subpopulation_idx), 1] X_initial_growth = hcat(rep_vector(1.0, num_elements(subpopulation_idx)));
int pop_initial_growth_n_covariates = 1;
matrix[num_elements(ww_lab_idx), 1] X_log_sigma_ww = hcat(rep_vector(1.0, num_elements(ww_lab_idx)));
int pop_log_sigma_ww_n_covariates = 1;
matrix[num_elements(count_stream_idx), 1] X_count_rate_center = hcat(rep_vector(1.0, num_elements(count_stream_idx)));
int pop_count_rate_center_n_covariates = 1;
matrix[num_elements(count_stream_idx), 1] X_log_phi_count = hcat(rep_vector(1.0, num_elements(count_stream_idx)));
int pop_log_phi_count_n_covariates = 1;
}
parameters {
cholesky_factor_corr[n_terms_initial_subpopulation] b_initial_subpopulation_L;
vector<lower=0.0>[n_terms_initial_subpopulation] b_initial_subpopulation_tau;
vector[(n_terms_initial_subpopulation * n_subpopulation)] b_initial_subpopulation_z_flat;
cholesky_factor_corr[n_terms_growth_subpopulation] b_growth_subpopulation_L;
vector<lower=0.0>[n_terms_growth_subpopulation] b_growth_subpopulation_tau;
vector[(n_terms_growth_subpopulation * n_subpopulation)] b_growth_subpopulation_z_flat;
cholesky_factor_corr[n_terms_ww_scale_ww_lab] b_ww_scale_ww_lab_L;
vector<lower=0.0>[n_terms_ww_scale_ww_lab] b_ww_scale_ww_lab_tau;
vector[(n_terms_ww_scale_ww_lab * n_ww_lab)] b_ww_scale_ww_lab_z_flat;
cholesky_factor_corr[n_terms_ww_noise_ww_lab] b_ww_noise_ww_lab_L;
vector<lower=0.0>[n_terms_ww_noise_ww_lab] b_ww_noise_ww_lab_tau;
vector[(n_terms_ww_noise_ww_lab * n_ww_lab)] b_ww_noise_ww_lab_z_flat;
cholesky_factor_corr[n_terms_count_rate_count_stream] b_count_rate_count_stream_L;
vector<lower=0.0>[n_terms_count_rate_count_stream] b_count_rate_count_stream_tau;
vector[(n_terms_count_rate_count_stream * n_count_stream)] b_count_rate_count_stream_z_flat;
cholesky_factor_corr[n_terms_count_dispersion_count_stream] b_count_dispersion_count_stream_L;
vector<lower=0.0>[n_terms_count_dispersion_count_stream] b_count_dispersion_count_stream_tau;
vector[(n_terms_count_dispersion_count_stream * n_count_stream)] b_count_dispersion_count_stream_z_flat;
real<lower=0.0> gamma;
real<lower=0, upper=1> phi_delta;
real<lower=0.0> sigma_delta;
real log10_g;
real<lower=0.0> t_peak;
real viral_peak;
real<lower=0.0> shed_tail;
vector[pop_log_ru_week_n_covariates] pop_log_ru_week_beta_pop;
real<lower=0.0, upper=1.0> dar_log_ru_week_week_grid_beta;
real<lower=0.0> dar_log_ru_week_week_grid_sigma;
vector[dar_log_ru_week_week_grid_n_innov] dar_log_ru_week_week_grid_z;
vector[pop_logit_I0_n_covariates] pop_logit_I0_beta_pop;
vector[pop_initial_growth_n_covariates] pop_initial_growth_beta_pop;
matrix[n_total, kernel_nsub_I_mat] I_mat_eps_d;
vector[pop_log_sigma_ww_n_covariates] pop_log_sigma_ww_beta_pop;
real<lower=0, upper=1> phi_count;
real<lower=0.0> sigma_count;
vector[pop_count_rate_center_n_covariates] pop_count_rate_center_beta_pop;
matrix[n_weeks, kernel_nsub_count_rate_week] count_rate_week_eps_count;
simplex[dow_share_alpha_n] dow_share;
vector[pop_log_phi_count_n_covariates] pop_log_phi_count_beta_pop;
}
transformed parameters {
matrix[n_terms_initial_subpopulation, n_subpopulation] b_initial_subpopulation_z = to_matrix(b_initial_subpopulation_z_flat, n_terms_initial_subpopulation, n_subpopulation);
matrix[n_subpopulation, n_terms_initial_subpopulation] b_initial_subpopulation = ((
diag_pre_multiply(b_initial_subpopulation_tau, b_initial_subpopulation_L) *
b_initial_subpopulation_z
)');
matrix[n_terms_growth_subpopulation, n_subpopulation] b_growth_subpopulation_z = to_matrix(b_growth_subpopulation_z_flat, n_terms_growth_subpopulation, n_subpopulation);
matrix[n_subpopulation, n_terms_growth_subpopulation] b_growth_subpopulation = ((diag_pre_multiply(b_growth_subpopulation_tau, b_growth_subpopulation_L) * b_growth_subpopulation_z)');
matrix[n_terms_ww_scale_ww_lab, n_ww_lab] b_ww_scale_ww_lab_z = to_matrix(b_ww_scale_ww_lab_z_flat, n_terms_ww_scale_ww_lab, n_ww_lab);
matrix[n_ww_lab, n_terms_ww_scale_ww_lab] b_ww_scale_ww_lab = ((diag_pre_multiply(b_ww_scale_ww_lab_tau, b_ww_scale_ww_lab_L) * b_ww_scale_ww_lab_z)');
matrix[n_terms_ww_noise_ww_lab, n_ww_lab] b_ww_noise_ww_lab_z = to_matrix(b_ww_noise_ww_lab_z_flat, n_terms_ww_noise_ww_lab, n_ww_lab);
matrix[n_ww_lab, n_terms_ww_noise_ww_lab] b_ww_noise_ww_lab = ((diag_pre_multiply(b_ww_noise_ww_lab_tau, b_ww_noise_ww_lab_L) * b_ww_noise_ww_lab_z)');
matrix[n_terms_count_rate_count_stream, n_count_stream] b_count_rate_count_stream_z = to_matrix(b_count_rate_count_stream_z_flat, n_terms_count_rate_count_stream, n_count_stream);
matrix[n_count_stream, n_terms_count_rate_count_stream] b_count_rate_count_stream = ((
diag_pre_multiply(b_count_rate_count_stream_tau, b_count_rate_count_stream_L) *
b_count_rate_count_stream_z
)');
matrix[n_terms_count_dispersion_count_stream, n_count_stream] b_count_dispersion_count_stream_z = to_matrix(
b_count_dispersion_count_stream_z_flat,
n_terms_count_dispersion_count_stream,
n_count_stream
);
matrix[n_count_stream, n_terms_count_dispersion_count_stream] b_count_dispersion_count_stream = ((
diag_pre_multiply(b_count_dispersion_count_stream_tau, b_count_dispersion_count_stream_L) *
b_count_dispersion_count_stream_z
)');
real dur_shed = (t_peak + shed_tail);
vector[nsh] shedding = viral_shedding_trajectory(t_peak, viral_peak, dur_shed, nsh);
vector[num_elements(week_grid)] pop_log_ru_week = (X_log_ru_week * pop_log_ru_week_beta_pop);
vector[(dar_log_ru_week_week_grid_n_innov + 1)] dar_log_ru_week_week_grid = differenced_ar1_path(
dar_log_ru_week_week_grid_beta,
dar_log_ru_week_week_grid_sigma,
dar_log_ru_week_week_grid_z
);
vector[num_elements(week_grid)] log_ru_week = (pop_log_ru_week + dar_log_ru_week_week_grid);
vector[week_idx_n] log_ru = weekly_expand(log_ru_week, week_idx);
vector[num_elements(subpopulation_idx)] pop_logit_I0 = (X_logit_I0 * pop_logit_I0_beta_pop);
vector[subpopulation_idx_n] r_logit_I0_initial_subpopulation = b_initial_subpopulation[subpopulation_idx, 1];
vector[num_elements(subpopulation_idx)] logit_I0 = (pop_logit_I0 + r_logit_I0_initial_subpopulation);
vector[num_elements(subpopulation_idx)] pop_initial_growth = (X_initial_growth * pop_initial_growth_beta_pop);
vector[subpopulation_idx_n] r_initial_growth_growth_subpopulation = b_growth_subpopulation[subpopulation_idx, 1];
vector[num_elements(subpopulation_idx)] initial_growth = (pop_initial_growth + r_initial_growth_growth_subpopulation);
matrix[n_total, kernel_nsub_I_mat] I_mat_delta_k;
matrix[week_idx_n, kernel_nsub_I_mat] I_mat_logRt_k;
matrix[week_idx_n, kernel_nsub_I_mat] I_mat;
for(plate_i__pl_1 in 1:kernel_nsub_I_mat) {
I_mat_delta_k[:, plate_i__pl_1] = ar1_dev(I_mat_eps_d[:, plate_i__pl_1], phi_delta, sigma_delta);
I_mat_logRt_k[:, plate_i__pl_1] = (log_ru + ((1.0 - is_reference[plate_i__pl_1]) * I_mat_delta_k[:, plate_i__pl_1]));
I_mat[:, plate_i__pl_1] = renewal_from_first_observed(
I_mat_logRt_k[:, plate_i__pl_1],
g,
gamma,
inv_logit(logit_I0[plate_i__pl_1]),
initial_growth[plate_i__pl_1],
uot
);
}
vector[ww_lab_idx_n] r_log_lab_mod_ww_scale_ww_lab = b_ww_scale_ww_lab[ww_lab_idx, 1];
vector[ww_lab_idx_n] log_lab_mod = r_log_lab_mod_ww_scale_ww_lab;
vector[num_elements(ww_lab_idx)] pop_log_sigma_ww = (X_log_sigma_ww * pop_log_sigma_ww_beta_pop);
vector[ww_lab_idx_n] r_log_sigma_ww_ww_noise_ww_lab = b_ww_noise_ww_lab[ww_lab_idx, 1];
vector[num_elements(ww_lab_idx)] log_sigma_ww = (pop_log_sigma_ww + r_log_sigma_ww_ww_noise_ww_lab);
vector[ww_lab_idx_n] ww_mu = ww_expected_log(I_mat, shedding, ww_time, ww_subpop, ww_lab_idx, log_lab_mod, log10_g, mwpd);
vector[ww_lab_idx_n] ww_sigma = gather_exp(log_sigma_ww, ww_lab_idx);
vector[num_elements(count_stream_idx)] pop_count_rate_center = (X_count_rate_center * pop_count_rate_center_beta_pop);
vector[count_stream_idx_n] r_count_rate_center_count_rate_count_stream = b_count_rate_count_stream[count_stream_idx, 1];
vector[num_elements(count_stream_idx)] count_rate_center = (pop_count_rate_center + r_count_rate_center_count_rate_count_stream);
matrix[n_weeks, kernel_nsub_count_rate_week] count_rate_week;
for(plate_i__pl_2 in 1:kernel_nsub_count_rate_week) {
count_rate_week[:, plate_i__pl_2] = (
count_rate_center[plate_i__pl_2] +
ar1_dev(count_rate_week_eps_count[:, plate_i__pl_2], phi_count, sigma_count)
);
}
matrix[week_idx_n, kernel_nsub_count_rate_week] count_rate_daily = weekly_expand_columns(count_rate_week, week_idx);
vector[dow_share_alpha_n] dow_effect = scale_simplex(dow_share, 7.0);
vector[num_elements(count_stream_idx)] pop_log_phi_count = (X_log_phi_count * pop_log_phi_count_beta_pop);
vector[count_stream_idx_n] r_log_phi_count_count_dispersion_count_stream = b_count_dispersion_count_stream[count_stream_idx, 1];
vector[num_elements(count_stream_idx)] log_phi_count = (pop_log_phi_count + r_log_phi_count_count_dispersion_count_stream);
vector[count_stream_idx_n] count_mu = count_interval_mean(
I_mat,
count_subpop_weights,
count_delay,
count_rate_daily,
count_start,
count_stop,
count_stream_idx,
dow,
dow_effect,
count_population
);
vector[count_stream_idx_n] count_phi = gather_exp(log_phi_count, count_stream_idx);
}
model {
b_initial_subpopulation_L ~ lkj_corr_cholesky(1.0);
b_initial_subpopulation_tau ~ normal(0.0, 0.75);
b_initial_subpopulation_z_flat ~ std_normal();
b_growth_subpopulation_L ~ lkj_corr_cholesky(1.0);
b_growth_subpopulation_tau ~ normal(0.0, 0.05);
b_growth_subpopulation_z_flat ~ std_normal();
b_ww_scale_ww_lab_L ~ lkj_corr_cholesky(1.0);
b_ww_scale_ww_lab_tau ~ normal(0.0, 0.5);
b_ww_scale_ww_lab_z_flat ~ std_normal();
b_ww_noise_ww_lab_L ~ lkj_corr_cholesky(1.0);
b_ww_noise_ww_lab_tau ~ normal(0.0, 0.35);
b_ww_noise_ww_lab_z_flat ~ std_normal();
b_count_rate_count_stream_L ~ lkj_corr_cholesky(1.0);
b_count_rate_count_stream_tau ~ normal(0.0, 0.5);
b_count_rate_count_stream_z_flat ~ std_normal();
b_count_dispersion_count_stream_L ~ lkj_corr_cholesky(1.0);
b_count_dispersion_count_stream_tau ~ normal(0.0, 0.5);
b_count_dispersion_count_stream_z_flat ~ std_normal();
gamma ~ lognormal(-4.0, 0.5);
phi_delta ~ beta(2.0, 8.0);
sigma_delta ~ exponential((1.0 ./ 4.0));
log10_g ~ normal(12.0, 1.0);
t_peak ~ lognormal(1.3862943611198906, 0.25);
viral_peak ~ normal(6.0, 1.0);
shed_tail ~ lognormal(2.6390573296152584, 0.3);
pop_log_ru_week_beta_pop ~ normal([0.0]', [0.5]');
dar_log_ru_week_week_grid_beta ~ normal(0.5, 0.2);
dar_log_ru_week_week_grid_sigma ~ normal(0.0, 0.2);
dar_log_ru_week_week_grid_z ~ std_normal();
pop_logit_I0_beta_pop ~ normal([-9.0]', [1.0]');
pop_initial_growth_beta_pop ~ normal([0.0]', [0.05]');
for(plate_i__pl_1 in 1:kernel_nsub_I_mat) {
I_mat_eps_d[:, plate_i__pl_1] ~ std_normal();
}
pop_log_sigma_ww_beta_pop ~ normal([-1.0498221244986778]', [0.5]');
ww_log_conc ~ lower_clamping_normal(ww_lod, ww_mu, ww_sigma);
phi_count ~ beta(2.0, 8.0);
sigma_count ~ exponential((1.0 ./ 20.0));
pop_count_rate_center_beta_pop ~ normal([-4.6]', [0.3]');
for(plate_i__pl_2 in 1:kernel_nsub_count_rate_week) {
count_rate_week_eps_count[:, plate_i__pl_2] ~ std_normal();
}
dow_share ~ dirichlet(dow_share_alpha);
pop_log_phi_count_beta_pop ~ normal([2.70805020110221]', [0.5]');
count ~ neg_binomial_2(count_mu, count_phi);
}
generated quantities {
vector[ww_log_conc_n] ww_log_conc_likelihood = lower_clamping_normal_lpdfs(ww_log_conc, ww_lod, ww_mu, ww_sigma);
vector[ww_log_conc_n] ww_log_conc_gen = lower_clamping_vector_normal_rng(ww_log_conc_n, ww_lod, ww_mu, ww_sigma);
vector[week_idx_n] I_agg = wsum(I_mat, w);
vector[count_n] count_likelihood = neg_binomial_2_lpmfs(count, count_mu, count_phi);
array[count_n] int count_gen = neg_binomial_2_int_rng(count_n, count_mu, count_phi);
vector[ht] forecast_infections = take_window(I_agg, (uot + nt + 1), ht);
vector[forecast_count_stream_idx_n] forecast_count_mean = count_interval_mean(
I_mat,
count_subpop_weights,
count_delay,
count_rate_daily,
forecast_count_start,
forecast_count_stop,
forecast_count_stream_idx,
dow,
dow_effect,
count_population
);
vector[forecast_ww_lab_idx_n] forecast_ww_log_mean = ww_expected_log(
I_mat,
shedding,
forecast_ww_time,
forecast_ww_subpop,
forecast_ww_lab_idx,
log_lab_mod,
log10_g,
mwpd
);
}Turing unsupported for this BRM example
BRM preparation: `effect(log_ru_week, Intercept)` names no linear predictor in this backend plan. Available predictors: logit_I0, initial_growth, I_mat, log_lab_mod, log_sigma_ww.The Turing pane is retained even though this structural kernel is outside the current Turing executor. Its construction error documents that backend boundary.
What matches, and what does not yet match
The example preserves the high-level causal order from CDC's model_definition.md:
Transmission: unadjusted reproduction numbers drive feedback-adjusted renewal incidence.
Subpopulations: trajectories vary around a shared process and aggregate by population weight.
Wastewater: shedding-convolved subpopulation incidence drives censored wastewater measurements.
Counts: delay-convolved mapped incidence drives weekday-adjusted interval counts.
The current fixture now includes the CDC model's distinct latent and observation axes: a 50-day unobserved period, an uncovered reference population, sparse and repeated site/lab/time records, record-specific detection limits, lab-level scale and noise effects, a normalized inferred triangular shedding trajectory, hierarchical first-observed incidence and growth with CDC's seeding back-calculation, bounded stationary subpopulation AR deviations, the weekly differenced-AR reference process, and a mean-one simplex weekday multiplier.
It does not yet reproduce these defining CDC details:
CDC's exact mean-reverting IHR parameterization;
the exact upstream hyperprior values supplied by CDC's R interface; or
CDC's component switches.
The executable upstream wwinference.stan is authoritative for those details. “Structural port” here means that the four components are coupled in the same causal order; it does not mean that their parameterizations, data axes, or priors are identical.
Relation to the StanBlocks form
The reproduction file also carries a pure-StanBlocks @slic companion, cdc_ww_inference_model. It uses the differenced-AR helper directly and broadcasts subpopulation parameters with plate. The @brm kernel term serves the same sampling role: parameters such as subpopulation innovations and initial incidence are introduced per cell, while @deffun owns only deterministic recurrence.
Both forms use the same differenced-AR global R recurrence. The companion shares the remaining simplifications listed above and is not presented as a numerical reference implementation.
BRM authoring gap removed by this port
This example exposed that the ordinary ar term could not represent CDC's differenced-AR weekly process. BRM snag brm-formula-diff-52808dea added the direct summand dar(week_grid; p=1): the term owns bounded persistence, innovation scale, and standardized innovations, while the formula intercept supplies the initial level. Replay and descriptor semantics now derive from that term as they do for other formula components. The example uses it directly; the remaining differences above are model-port scope rather than this former authoring limitation.
Composable counts and forecasts
The count likelihood is not tied to a column called hosp or to one row per day. The fixture carries one row per observed interval:
count_start,count_stop, andcount_stream_idxdefine the record axis;count_subpop_weights[:, stream]maps the latent subpopulation matrix into that stream's catchment;count_delay[:, stream]supplies its infection-to-observation delay; andcount_population[stream]supplies the absolute-count scale.
The default fixture deliberately crosses two shapes: daily jurisdiction hospital admissions and weekly totals over only the sampled catchments. Adding another count source is therefore a data extension—append a stream column and its records—rather than a second hand-written likelihood.
Forecast mappings use the same deterministic functions but are not responses in the calibration likelihood. The emitted Stan descriptor owns three named generated quantities: forecast_infections, forecast_count_mean, and forecast_ww_log_mean. The reproduction exposes both layers without a parallel registry:
plan = cdc_ww_brm_plan(df) # declarations and replay provenance
descriptor = cdc_ww_brm_descriptor(df)The descriptor derives fit, predict, pointwise_loglik, replay, and reprocess from the executable model. Consumers should resolve the named forecast outputs from descriptor.outputs; they should not parse generated Stan identifiers or infer them from declaration order.
Provenance
Reproduction: research/wastewater/cdc_ww_inference.jl, gated by test/cdc_ww_inference.jl. Both the @slic and @brm forms are checked for transpilation and stanc; the default fixtures are also checked for a finite BridgeStan density and gradient. The discretized kernels and priors remain illustrative. The models are not sampled here and the gate is not a posterior-equivalence test against CDC.
Run the reproduction
After bootstrapping the repository's test environment:
julia --startup-file=no --project=test test/cdc_ww_inference.jlSet BRM_KERNEL_RUNTIME=0 to run lowering and stanc without BridgeStan instantiation.