Skip to content

Renewal model with negative-binomial reporting ​

The renewal equation expresses new infections as a function of past infections weighted by the generation interval, scaled by a time-varying reproduction number [1]. Mishra et al. [2] showed that this construction follows from an age-dependent branching process and pairs with a negative-binomial observation model to give a Bayesian hierarchical model for reported case counts.

This tutorial builds that model from two composed parts and fits it to the test-confirmed COVID-19 cases from South Korea that Mishra et al. [2] analysed. The parts are a Renewal infection process carrying an autoregressive latent process for , and a NegativeBinomialError observation model. The latent process is folded into the renewal model rather than supplied as a separate top-level component, because the reproduction number is the renewal model's own parameter process.

The model ​

is the discretised generation interval, the autoregressive damping, the innovation standard deviation, and the observation overdispersion.

Components ​

The latent process is a second-order autoregressive model on with a HierarchicalNormal innovation term, matching Mishra et al. [2]. Strong autocorrelation in the reproduction number is encoded by a first damping prior concentrated near one,   on , with a weaker second lag.

julia
using ComposableTuringIDModels, Distributions, Random, Turing, Mooncake
using ADTypes: AutoMooncake
Random.seed!(1234)

latent = AR(
    damp = [truncated(Normal(0.8, 0.05), 0, 1),
        truncated(Normal(0.1, 0.05), 0, 1)],
    init = [Normal(0.0, 0.2), Normal(0.0, 0.2)],
    ϵ_t = HierarchicalNormal(std = HalfNormal(0.1)))
AR
└─ ϵ_t: HierarchicalNormal

The infection process needs a discrete generation interval. Renewal takes a continuous distribution and discretises it with double interval censoring [9], using CensoredDistributions.jl. Following Mishra et al. [2] we use a serial interval as a proxy for the generation interval. Renewal is the only infection model that carries a generation interval, because it is the only one that uses one. It couples that interval to the latent process in its rt slot and a prior for the initial infections.

julia
renewal = Renewal(; generation_time = Gamma(6.5, 0.62),
    rt = latent, initialisation = Normal(log(1.0), 0.1))
renewal.gen_int
8-element Vector{Float64}:
 0.026663134095601098
 0.14059778064943768
 0.2502660305615845
 0.24789569560506872
 0.1731751163417783
 0.09635404000022221
 0.045734375752163825
 0.01931382699414364

The stored gen_int is a probability vector, the continuous serial interval binned into daily weights that sum to one. Double interval censoring is not the same as evaluating the continuous density at integer days. It accounts for both the primary and secondary events falling anywhere within their days [9]. The lag-0 bin is dropped and the weights renormalised over the days kept, so the discrete interval is a little shorter and a little tighter than the behind it.

julia
using Statistics
lags = 1:length(renewal.gen_int)
discrete_mean = sum(lags .* renewal.gen_int)
discrete_sd = sqrt(sum((lags .- discrete_mean) .^ 2 .* renewal.gen_int))
serial_interval = Gamma(6.5, 0.62)
(discrete = round.((discrete_mean, discrete_sd), digits = 2),
    continuous = round.((mean(serial_interval), std(serial_interval)), digits = 2))
(discrete = (3.97, 1.53), continuous = (4.03, 1.58))

The infection process in isolation ​

The renewal model is a model in its own right, so it can be exercised without an observation model. Pinning its reproduction-number process to a known trajectory isolates the contribution of the renewal equation. With the latent folded in, that means building a renewal model whose rt slot is a deterministic FixedIntercept latent, giving a constant , and fixing the initial-infections parameter. The same as_turing_model call that composes into the full model then runs the infection process standalone, returning its infections I_t and the internal latent draw Z_t.

julia
fixed_logR = log(1.4)
renewal_fixed = Renewal(; generation_time = renewal.gen_int,
    rt = FixedIntercept(fixed_logR), initialisation = Normal())
demo = fix(as_turing_model(renewal_fixed, 60), (init_incidence = 0.0,))()
(constant_Rt = round(exp(first(demo.Z_t)), digits = 2),
    grows = demo.I_t[end] > demo.I_t[1])
(constant_Rt = 1.4, grows = true)

A constant   grows incidence. A path that declined through zero would instead produce the textbook turn-over, with incidence growing, decelerating as  , and falling once  . Driving the renewal model with a richer fixed path means swapping the FixedIntercept latent for a deterministic latent of the desired shape.

Reported cases are overdispersed counts of the latent infections. The prior is placed on the cluster factor , which is roughly the coefficient of variation of the observation noise and so easier to reason about a priori.

julia
obs = NegativeBinomialError(cluster_factor = HalfNormal(0.1))

The renewal infection process already carries the latent process, so IDModel needs only it and the observation model.

julia
model = IDModel(renewal, obs)
IDModel
├─ infection: Renewal
│  └─ rt: AR
│     └─ ϵ_t: HierarchicalNormal
└─ observation: NegativeBinomialError

The composed model is also a prior simulator. Passing missing observations makes as_turing_model return generated quantities instead of conditioning on data: the reported cases generated_y_t, the latent infections I_t, and the latent process Z_t. Here we go straight to real data.

The data ​

Mishra et al. [2] fit this model to daily test-confirmed COVID-19 cases in South Korea over the first wave of 2020. The series is stored with the docs and read with CSV and DataFrames.

julia
using CSV, DataFrames
datapath = joinpath(pkgdir(ComposableTuringIDModels),
    "docs", "src", "tutorials", "data", "south_korea_data.csv")
south_korea = CSV.read(datapath, DataFrame)
first(south_korea, 5)
5×4 DataFrame
RowColumn1datecases_newdeaths_new
Int64DateInt64Int64
112019-12-3100
222020-01-0100
332020-01-0200
442020-01-0300
552020-01-0400

We fit the growth-and-decline window of the first wave, matching the span used by Mishra et al. [2], and take the reported cases over it as the observed series.

julia
tspan = (45, 80)
y_obs = south_korea.cases_new[first(tspan):last(tspan)]
n = length(y_obs)
(n = n, total_cases = sum(y_obs),
    from = south_korea.date[first(tspan)], to = south_korea.date[last(tspan)])
(n = 36, total_cases = 8537, from = Dates.Date("2020-02-13"), to = Dates.Date("2020-03-19"))

Fit ​

Conditioning on the observed counts and sampling with NUTS recovers the posterior. We draw two chains in parallel with MCMCThreads() so the posterior is well resolved and the cross-chain diagnostic is available. The slightly raised target acceptance rate keeps the sampler stable on the hierarchical innovation scale. We differentiate with Mooncake, the recommended backend for this package, described under Automatic differentiation backend.

julia
posterior = as_turing_model(model, y_obs, n)
chain = sample(
    posterior, NUTS(0.95; adtype = AutoMooncake(; config = nothing)),
    MCMCThreads(), 250, 2; progress = false)
┌ Info: Found initial step size
└   ϵ = 0.0015625
┌ Info: Found initial step size
└   ϵ = 0.003125
┌ Warning: There were 8 divergent transitions. Consider reparameterising your model or using a smaller step size. For adaptive samplers such as NUTS and HMCDA, consider increasing `target_accept`.
└ @ Turing.Inference ~/.julia/packages/Turing/4sXd9/src/mcmc/hmc.jl:483
┌ Warning: There were 8 divergent transitions. Consider reparameterising your model or using a smaller step size. For adaptive samplers such as NUTS and HMCDA, consider increasing `target_accept`.
└ @ Turing.Inference ~/.julia/packages/Turing/4sXd9/src/mcmc/hmc.jl:483

Sampling returns a chain whose parameter names are flat. Composition does not namespace by default, so a component's own parameter keeps its plain name however deeply it is nested, and damp and std below belong to the AR process folded into the renewal model. A few components deliberately prefix their children, and two components that would otherwise use the same name are separated by wrapping one in PrefixLatentModel or PrefixObservationModel. sample returns a FlexiChains chain, which summarystats summarises directly.

julia
using MCMCChains
summarystats(chain)
╭─FlexiSummary (9 statistics) ─────────────────────────────────────────────────╮
│   iter    collapsed                                                          │
│   chain   collapsed                                                          │
│ ↓ stat  = [mean, std, mcse, ess_bulk, ess_tail, rhat, q5, q50, q95]          │
│                                                                              │
│ Parameters (41) ── AbstractPPL.VarName                                       │
│  Float64  init[1], init[2], damp[1], damp[2], std, ϵ_t[1], ϵ_t[2], ϵ_t[3],   │
│           ϵ_t[4], ϵ_t[5], ϵ_t[6], ϵ_t[7], ϵ_t[8], ϵ_t[9], ϵ_t[10], ϵ_t[11],  │
│           ϵ_t[12], ϵ_t[13], ϵ_t[14], ϵ_t[15], ϵ_t[16], ϵ_t[17], ϵ_t[18],     │
│           ϵ_t[19], ϵ_t[20], ϵ_t[21], ϵ_t[22], ϵ_t[23], ϵ_t[24], ϵ_t[25],     │
│           ϵ_t[26], ϵ_t[27], ϵ_t[28], ϵ_t[29], ϵ_t[30], ϵ_t[31], ϵ_t[32],     │
│           ϵ_t[33], ϵ_t[34], init_incidence, cluster_factor                   │
│                                                                              │
│ Extras (14)                                                                  │
│  Float64  n_steps, is_accept, acceptance_rate, log_density,                  │
│           hamiltonian_energy, hamiltonian_energy_error,                      │
│           max_hamiltonian_energy_error, tree_depth, numerical_error,         │
│           step_size, nom_step_size, logprior, loglikelihood, logjoint        │
│                                                                              │
│ Summary                                                                      │
│          param     mean     std    mcse   ess_bulk  ess_tail    rhat  …      │
│        init[1]   0.0346  0.1950  0.0079   588.2413  407.5449  0.9983  …      │
│        init[2]  -0.0093  0.1958  0.0079   613.8495  375.4466  0.9994  …      │
│        damp[1]   0.8178  0.0395  0.0023   291.7387  352.1886  1.0009  …      │
│        damp[2]   0.0822  0.0383  0.0031   147.0113  122.9419  0.9996  …      │
│            std   0.4236  0.0461  0.0031   221.2322  336.2500  1.0055  …      │
│         ϵ_t[1]   0.5796  0.8643  0.0866   104.1552  304.9701  1.0049  …      │
│         ϵ_t[2]   0.8265  0.9273  0.0912   103.1941  168.9442  1.0039  …      │
│         ϵ_t[3]   1.1432  0.8974  0.0794   134.2592  187.4583  1.0025  …      │
│         ϵ_t[4]   1.3226  0.8432  0.0661   162.1429  173.2636  1.0068  …      │
│         ϵ_t[5]   2.4027  0.7327  0.0602   148.9410  297.9942  1.0046  …      │
│         ϵ_t[6]   1.5284  0.6472  0.0431   226.3204  347.8869  1.0068  …      │
│         ϵ_t[7]   0.8821  0.5939  0.0381   244.1199  330.8633  0.9984  …      │
│         ϵ_t[8]   0.7140  0.5163  0.0367   196.6669  243.5891  1.0015  …      │
│         ϵ_t[9]  -0.6227  0.4778  0.0275   306.0629  276.0184  1.0014  …      │
│        ϵ_t[10]  -2.2873  0.5022  0.0297   290.4668  198.7462  1.0032  …      │
│        ϵ_t[11]  -1.6785  0.4575  0.0330   209.6017  358.7342  1.0001  …      │
│        ϵ_t[12]   0.5643  0.4031  0.0234   297.5382  217.6172  0.9985  …      │
│        ϵ_t[13]   1.1107  0.3711  0.0216   315.8047  316.0487  0.9996  …      │
│        ϵ_t[14]   0.0094  0.3134  0.0134   532.2502  310.2869  1.0005  …      │
│        ϵ_t[15]   1.2674  0.3686  0.0248   234.9056  178.0394  1.0027  …      │
│        ϵ_t[16]  -1.1212  0.3935  0.0222   321.0003  339.7730  1.0007  …      │
│        ϵ_t[17]  -0.4229  0.3080  0.0157   386.5337  261.7668  1.0002  …      │
│        ϵ_t[18]  -0.7732  0.3232  0.0158   440.5482  301.4940  1.0023  …      │
│        ϵ_t[19]  -0.6797  0.3474  0.0125   775.9278  317.3573  1.0198  …      │
│        ϵ_t[20]  -0.5633  0.3370  0.0157   485.0482  170.1333  1.0074  …      │
│        ϵ_t[21]   0.2557  0.3480  0.0184   391.7184  129.8265  1.0107  …      │
│        ϵ_t[22]  -0.0282  0.3191  0.0113   907.0854  125.5653  1.0034  …      │
│        ϵ_t[23]  -0.5333  0.3157  0.0100  1126.4719  149.8924  1.0151  …      │
│        ϵ_t[24]  -0.8986  0.3655  0.0160   539.9142  257.4254  0.9997  …      │
│        ϵ_t[25]  -1.3419  0.3975  0.0195   467.5143  242.1348  1.0022  …      │
│        ϵ_t[26]   1.0270  0.4083  0.0345   146.7370  269.5487  1.0197  …      │
│        ϵ_t[27]  -1.1345  0.4129  0.0315   185.8281  181.3857  1.0188  …      │
│        ϵ_t[28]   0.0400  0.4046  0.0212   390.6716  184.1683  1.0187  …      │
│        ϵ_t[29]   0.1846  0.3872  0.0207   321.5629  414.1465  0.9987  …      │
│        ϵ_t[30]  -0.3170  0.4533  0.0218   434.2766  339.3163  1.0008  …      │
│        ϵ_t[31]   0.1494  0.4243  0.0212   410.5655  223.0201  0.9993  …      │
│        ϵ_t[32]   0.6283  0.4259  0.0275   245.3914  322.2411  1.0104  …      │
│        ϵ_t[33]   0.6238  0.4166  0.0182   516.7507  408.5027  1.0200  …      │
│        ϵ_t[34]   1.2810  0.3800  0.0239   267.6385  141.7101  0.9982  …      │
│   init_incide…  -0.0401  0.0980  0.0044   480.6895  388.8842  1.0158  …      │
│   cluster_fac…   0.0868  0.0556  0.0077    44.3583   81.9428  1.0549  …      │
╰──────────────────────────────────────────────────────────────────────────────╯

Prior versus posterior ​

Sampling the same model with Prior, ignoring the observations, gives a prior draw over the same parameters. Overlaying it on the posterior shows which parameters moved. We load a Makie backend and PairPlots.jl. The FlexiChains PairPlots extension turns a chain subset to a few keys with chain[[...]] into a PairPlots.Series, so prior and posterior overlay on one corner plot.

julia
using CairoMakie, PairPlots

prior_chain = sample(posterior, Prior(), 1000; progress = false)
pp_keys = [@varname(damp), @varname(std),
    @varname(cluster_factor), @varname(init_incidence)]
pairplot(
    PairPlots.Series(chain[pp_keys]; label = "posterior"),
    PairPlots.Series(prior_chain[pp_keys]; label = "prior"))

The innovation scale (std) is sharply updated away from its prior, so the data are informative about how much moves. The autoregressive damping (damp), the cluster factor and the initial infections stay closer to their priors on this short window.

Posterior trajectories ​

The reproduction number   is a generated quantity rather than a sampled parameter. generated_observables re-runs the fitted model over the chain to recover the latent and infection trajectories per draw. The reported counts are scored element-wise, so their posterior predictive distribution comes from predict on the same model with the observations set to missing.

A couple of small helpers reduce the per-draw trajectories to credible bands and draw a median line with 50% and 95% ribbons.

Stack the per-draw into an band, draw the posterior-predictive from the unconditioned model, and plot both against the observed series.

julia
gens = vec(generated_observables(posterior, y_obs, chain).generated)
Rt = credible_bands(reduce(hcat, (exp.(g.Z_t) for g in gens)))

pred = predict(as_turing_model(model, fill(missing, n), n), chain)
yt = predictive_bands(pred, n)

fig = Figure(; size = (760, 620))
ax1 = Axis(fig[1, 1]; ylabel = "Reproduction number Rₜ")
ci_ribbon!(ax1, 1:size(Rt, 1), Rt; color = :purple, label = "posterior median")
hlines!(ax1, [1.0]; color = :grey, linestyle = :dash)
axislegend(ax1; position = :rt)
ax2 = Axis(fig[2, 1]; xlabel = "Day", ylabel = "Reported cases")
ci_ribbon!(ax2, 1:size(yt, 1), yt; color = :teal,
    label = "posterior predictive")
scatter!(ax2, 1:n, y_obs; color = :black, markersize = 7, label = "observed")
axislegend(ax2; position = :lt)
fig

The posterior-predictive band tracks the observed South Korean series closely. The path recovers the first-wave turn-over, with an early rise well above one, a fall through   as the wave peaks, and a decline below one as cases drop.

Forecasting the next weeks ​

The same fitted model forecasts out of sample in one call. The latent AR process is non-centred, so forecast carries each posterior draw forward and then draws the future reported cases. It holds the fitted parameters and the in-sample path fixed, and continues the process over the horizon with fresh prior innovations.

julia
h = 14
fc = forecast(model, y_obs, chain, h)
round.(quantile(Float64.(vec(fc[@varname(y_t[n + h])])), [0.05, 0.5, 0.95]),
    digits = 1)
3-element Vector{Float64}:
     6.0
   302.5
 50548.4

The returned chain carries the predicted over    .

julia
# Multi-level CI band quantiles: 90%, 60%, 30% + median
FC_CI_QS = [0.05, 0.2, 0.35, 0.5, 0.65, 0.8, 0.95]

# Draw three nested CI ribbons with the median line
function multi_ci_ribbon!(ax, ts, bands; color, label)
    # 90% CI (cols 1, 7)
    band!(ax, ts, bands[:, 1], bands[:, 7]; color = (color, 0.1))
    # 60% CI (cols 2, 6)
    band!(ax, ts, bands[:, 2], bands[:, 6]; color = (color, 0.25))
    # 30% CI (cols 3, 5)
    band!(ax, ts, bands[:, 3], bands[:, 5]; color = (color, 0.5))
    # median (col 4)
    lines!(ax, ts, bands[:, 4]; color = color, linewidth = 2, label = label)
end

# Multi-level credible bands for the forecast
fc_bands = credible_bands(reduce(vcat,
    (permutedims(vec(fc[@varname(y_t[i])])) for i in (n + 1):(n + h)));
    qs = FC_CI_QS)

# Sample 100 random forecast trajectories
n_draws = length(vec(fc[@varname(y_t[n + 1])]))
ntraj = min(100, n_draws)
sample_idx = rand(1:n_draws, ntraj)
trajectories = reduce(hcat,
    [[vec(fc[@varname(y_t[i])])[idx] for i in (n + 1):(n + h)] for idx in sample_idx])

fig_fc = Figure(; size = (760, 360))
axf = Axis(fig_fc[1, 1]; xlabel = "Day", ylabel = "Reported cases",
    yscale = log10)
scatter!(axf, 1:n, y_obs; color = :black, markersize = 7, label = "observed")
# Faint individual forecast trajectories
for idx in 1:ntraj
    lines!(axf, (n + 1):(n + h), max.(trajectories[:, idx], 1);
        color = (:teal, 0.08), linewidth = 0.5)
end
# Multi-level CI bands (offset to keep log10 scale safe at zero)
multi_ci_ribbon!(axf, (n + 1):(n + h), max.(fc_bands, 1); color = :teal,
    label = "forecast")
vlines!(axf, [n + 0.5]; color = :grey, linestyle = :dash)
axislegend(axf; position = :lt)
fig_fc

Swap a component ​

The parts share one interface, so an alternative observation assumption is a one-line change. Swapping the negative-binomial reporting for a PoissonError leaves the renewal infection process and its latent process untouched.

julia
using ComposableTuringIDModels: swap
poisson_model = swap(
    err -> err isa NegativeBinomialError ? PoissonError() : err, model)
IDModel
├─ infection: Renewal
│  └─ rt: AR
│     └─ ϵ_t: HierarchicalNormal
└─ observation: PoissonError

swap walks the observation chain, replaces every component matching the predicate and rebuilds each wrapper around its replacement. See Inspecting and updating an observation chain for the general form, which also targets a single stream of a Split.

References ​

  1. A. Cori, N. M. Ferguson, C. Fraser and S. Cauchemez. A new framework and software to estimate time-varying reproduction numbers during epidemics. American Journal of Epidemiology 178, 1505–1512 (2013).

  2. S. Mishra, T. Berah, T. A. Mellan, H. J. Unwin, M. A. Vollmer, K. V. Parag, A. Gandy, S. Flaxman and S. Bhatt. On the derivation of the renewal equation from an age-dependent branching process: an epidemic modelling perspective, arXiv preprint arXiv:2006.16487 (2020).

  3. K. Charniga, S. W. Park, A. R. Akhmetzhanov, A. Cori, J. Dushoff, S. Funk and others. Best practices for estimating and reporting epidemiological delay distributions of infectious diseases. PLoS Computational Biology 20, e1012520 (2024).