Skip to content

Continuous-time state-space models ​

A panel of people answers a few questions on their phones, several times a day, at irregular times. Behind the answers are latent processes — stress, mood — that evolve continuously, push on each other, and differ between people. This page builds that kind of model in BRM four ways, from sample every latent state to integrate every latent state out, fits it, and checks how much the numerical approximation inside the likelihood can be trusted.

The models are the demonstration models of ctsem, Charles Driver's R package for hierarchical continuous-time dynamic modelling (Driver, Oud & Voelkle 2017; Driver & Voelkle 2018). The model specifications, their parameterisation and the generating values used below are his; this page is about expressing and fitting them with @brm. Everything shown is built or measured from the sources in research/ema_ctsem/; all data are simulated.

The model class ​

Each subject i has a latent state ηi(t) (here: stress and mood) following a stochastic differential equation, observed at irregular times ti1<ti2<…:

dηi(t)=(Ai(ηi)ηi(t)+bi)dt+Gi(ηi,xi)dW(t),latent dynamicsηi(tij+)=ηi(tij−)+Mixij,input impulses at observationsyij∼p(⋅∣Ληi(tij)+τ),measurement

with Gaussian measurement error for continuous indicators and a logistic link for binary ones. The drift A and the diffusion G may depend on inputs and on the latent state itself; the subject-level parameters in Ai,bi,Gi,Mi follow a population model with covariates and correlated random effects.

Three layers, one model ​

Every model on this page divides the work the same way:

layerholdswritten as
formula surfacethe population model: covariates and (correlated) random effects on subject-level parametersb0 ~ 1 + age + treatment + (1 | p | subject)
kernel(...) cellone subject: that subject's series and that subject's parameter valuespred ~ kernel(dt, stressReport, …, b0, …) do … end
@deffunthe recurrence — a loop with carried state, emitted as a Stan functionan SDE path, or a filter

kernel(...) lowers to a loop over subjects whose body is the cell. What differs between the four models is only what the cell does with the latent path: sample it, or integrate it out. The population model on the formula surface is untouched by that choice.

1. Sampling the latent states ​

The direct translation: the latent path is a deterministic function of standard-normal innovations, and the innovations are parameters of the cell. The recurrence is an Euler–Maruyama scan:

julia
StanBlocks.@deffun begin
    # Euler-Maruyama path of the coupled SDE, driven by the innovations zs, zm.
    # Returns [stress; mood] stacked.
    ema_path(dt::vector[nt], wl::vector[nt],
             b0::real, bm::real, a12::real, a21::real, a22::real,
             cm::real, wls::real, q0::real, qw::real, diffm::real, diff21::real,
             s0::real, m0::real, zs::vector[nt], zm::vector[nt])::vector[2 * nt] = begin
        out::vector[2 * nt]
        s = s0 + wls * wl[1]; m = m0           # workload acts as an impulse at each observation
        corr = tanh(diff21)
        out[1] = s; out[nt + 1] = m
        for t in 2:nt
            drift_s = -log1p(exp(b0 + bm * m)) * s + a12 * m
            drift_m = a21 * s + a22 * m + cm
            gs = exp(q0 + qw * wl[t])              # stress volatility depends on the workload input
            zc = corr * zs[t] + sqrt(1 - corr * corr) * zm[t]
            s = s + drift_s * dt[t] + gs * sqrt(dt[t]) * zs[t] + wls * wl[t]
            m = m + drift_m * dt[t] + diffm * sqrt(dt[t]) * zc
            out[t] = s; out[nt + t] = m
        end
        out
    end
end

The model: stress recovers at a rate that depends on mood (a softplus keeps it negative), stress volatility depends on the workload input, the two shocks are correlated, workload pushes stress by an impulse at each observation, and three indicators — two continuous reports and a binary smoked — measure the states. Four subject-level parameters share one correlated random-effect block through the brms-style (1 | p | subject).

brm-comparison
Hierarchical EMA model, latent states sampled
julia
function ema_sampled_model(data = ema_sampled_fixture())
    @brm data begin
        sigma_s ~ Exponential(1); sigma_m ~ Exponential(1)          # measurement sds
        bm ~ Normal(0, 0.5); a12 ~ Normal(0, 0.5); a21 ~ Normal(0, 0.5); a22 ~ Normal(-0.5, 0.3)
        qw ~ Normal(0, 0.5); diffm ~ Exponential(1); diff21 ~ Normal(0, 0.5)
        l31 ~ Normal(0, 1); smoke_threshold ~ Normal(0, 1)          # binary indicator: loading, threshold
        mm_s ~ Normal(0, 0.5); mm_m ~ Normal(0, 0.5)                # manifest means
        # subject-level parameters: covariates + ONE correlated random-effect block
        b0  ~ 1 + age + treatment + (1 | p | subject)
        q0  ~ 1 +                   (1 | p | subject)
        cm  ~ 1 + treatment +       (1 | p | subject)
        wls ~ 1 +                   (1 | p | subject)
        s0 ~ 1 + (1 | subject)
        m0 ~ 1 + (1 | subject)
        stress ~ kernel(dt_grid, workload, stressReport, moodReport, smoked,
                        b0, q0, cm, wls, s0, m0) do dt, wl, sR, mR, smk, lb0, lq0, lcm, lwls, ls0, lm0
            zs::vector[dims(dt)[1]] ~ std_normal()                  # the innovations ARE parameters
            zm::vector[dims(dt)[1]] ~ std_normal()
            path = ema_path(dt, wl, lb0, bm, a12, a21, a22, lcm, lwls, lq0, qw, diffm, diff21, ls0, lm0, zs, zm)
            st = path[1:dims(dt)[1]]
            mo = path[(dims(dt)[1] + 1):(2 * dims(dt)[1])]
            sR ~ normal(mm_s .+ st, sigma_s)
            mR ~ normal(mm_m .+ mo, sigma_m)
            smk ~ bernoulli_logit(l31 .* st .+ smoke_threshold)
            st
        end
    end
end
julia
BRMI:
  sigma_s ~ Exponential(1)
  sigma_m ~ Exponential(1)
  bm ~ Normal(0, 0.5)
  a12 ~ Normal(0, 0.5)
  a21 ~ Normal(0, 0.5)
  a22 ~ Normal(-0.5, 0.3)
  qw ~ Normal(0, 0.5)
  diffm ~ Exponential(1)
  diff21 ~ Normal(0, 0.5)
  l31 ~ Normal(0, 1)
  smoke_threshold ~ Normal(0, 1)
  mm_s ~ Normal(0, 0.5)
  mm_m ~ Normal(0, 0.5)
  age: data (eltype=Float64, n=6)
  treatment: data (eltype=Float64, n=6)
  subject: data (eltype=String, n=6)
  b0 ~ 1 + age + treatment + (1 | p | subject)
  q0 ~ 1 + (1 | p | subject)
  cm ~ 1 + treatment + (1 | p | subject)
  wls ~ 1 + (1 | p | subject)
  s0 ~ 1 + (1 | subject)
  m0 ~ 1 + (1 | subject)
  dt_grid: data (eltype=Vector{Float64}, n=6)
  workload: data (eltype=Vector{Float64}, n=6)
  stressReport: data (eltype=Vector{Float64}, n=6)
  moodReport: data (eltype=Vector{Float64}, n=6)
  smoked: data (eltype=Vector{Int64}, n=6)
  stress ~ kernel((dt, wl, sR, mR, smk, lb0, lq0, lcm, lwls, ls0, lm0)->begin
        #= brm-docs-example.jl:17 =#
        zs::vector[(dims(dt))[1]] ~ std_normal()
        #= brm-docs-example.jl:18 =#
        zm::vector[(dims(dt))[1]] ~ std_normal()
        #= brm-docs-example.jl:19 =#
        path = ema_path(dt, wl, lb0, bm, a12, a21, a22, lcm, lwls, lq0, qw, diffm, diff21, ls0, lm0, zs, zm)
        #= brm-docs-example.jl:20 =#
        st = path[1:(dims(dt))[1]]
        #= brm-docs-example.jl:21 =#
        mo = path[(dims(dt))[1] + 1:2 * (dims(dt))[1]]
        #= brm-docs-example.jl:22 =#
        sR ~ normal(mm_s .+ st, sigma_s)
        #= brm-docs-example.jl:23 =#
        mR ~ normal(mm_m .+ mo, sigma_m)
        #= brm-docs-example.jl:24 =#
        smk ~ bernoulli_logit(l31 .* st .+ smoke_threshold)
        #= brm-docs-example.jl:25 =#
        st
    end, dt_grid, workload, stressReport, moodReport, smoked, b0, q0, cm, wls, s0, m0)
julia
SBBRMI with data keys = [:age, :dt_grid, :kernel_nsub_stress, :moodReport, :n_subject, :n_terms_p_subject, :smoked, :stressReport, :subject_idx, :treatment, :workload]
emitted @slic body:
begin
    b_p_subject ~ ranef_correlated_draws(; group_idx = subject_idx, n_groups = n_subject, n_terms = n_terms_p_subject)
    sigma_s ~ exponential(1.0 ./ 1)
    sigma_m ~ exponential(1.0 ./ 1)
    bm ~ normal(0, 0.5)
    a12 ~ normal(0, 0.5)
    a21 ~ normal(0, 0.5)
    a22 ~ normal(-0.5, 0.3)
    qw ~ normal(0, 0.5)
    diffm ~ exponential(1.0 ./ 1)
    diff21 ~ normal(0, 0.5)
    l31 ~ normal(0, 1)
    smoke_threshold ~ normal(0, 1)
    mm_s ~ normal(0, 0.5)
    mm_m ~ normal(0, 0.5)
    X_b0 = hcat(rep_vector(1.0, num_elements(age)), age, treatment)
    pop_b0 ~ popefs(; X = X_b0)
    r_b0_p_subject = b_p_subject[subject_idx, 1]
    b0 = pop_b0 + r_b0_p_subject
    X_q0 = hcat(rep_vector(1.0, num_elements(subject_idx)))
    pop_q0 ~ popefs(; X = X_q0)
    r_q0_p_subject = b_p_subject[subject_idx, 2]
    q0 = pop_q0 + r_q0_p_subject
    X_cm = hcat(rep_vector(1.0, num_elements(treatment)), treatment)
    pop_cm ~ popefs(; X = X_cm)
    r_cm_p_subject = b_p_subject[subject_idx, 3]
    cm = pop_cm + r_cm_p_subject
    X_wls = hcat(rep_vector(1.0, num_elements(subject_idx)))
    pop_wls ~ popefs(; X = X_wls)
    r_wls_p_subject = b_p_subject[subject_idx, 4]
    wls = pop_wls + r_wls_p_subject
    X_s0 = hcat(rep_vector(1.0, num_elements(subject_idx)))
    pop_s0 ~ popefs(; X = X_s0)
    r_s0_subject ~ ranef_intercept(; group_idx = subject_idx, n_groups = n_subject)
    s0 = pop_s0 + r_s0_subject
    X_m0 = hcat(rep_vector(1.0, num_elements(subject_idx)))
    pop_m0 ~ popefs(; X = X_m0)
    r_m0_subject ~ ranef_intercept(; group_idx = subject_idx, n_groups = n_subject)
    m0 = pop_m0 + r_m0_subject
    stress ~ plate(dt_grid, workload, stressReport, moodReport, smoked, b0, q0, cm, wls, s0, m0; outer = (kernel_nsub_stress,)) do dt, wl, sR, mR, smk, lb0, lq0, lcm, lwls, ls0, lm0
            #= brm-docs-example.jl:17 =#
            zs::vector[(dims(dt))[1]] ~ std_normal()
            #= brm-docs-example.jl:18 =#
            zm::vector[(dims(dt))[1]] ~ std_normal()
            #= brm-docs-example.jl:19 =#
            path = ema_path(dt, wl, lb0, bm, a12, a21, a22, lcm, lwls, lq0, qw, diffm, diff21, ls0, lm0, zs, zm)
            #= brm-docs-example.jl:20 =#
            st = path[1:(dims(dt))[1]]
            #= brm-docs-example.jl:21 =#
            mo = path[(dims(dt))[1] + 1:2 * (dims(dt))[1]]
            #= brm-docs-example.jl:22 =#
            sR ~ normal(mm_s .+ st, sigma_s)
            #= brm-docs-example.jl:23 =#
            mR ~ normal(mm_m .+ mo, sigma_m)
            #= brm-docs-example.jl:24 =#
            smk ~ bernoulli_logit(l31 .* st .+ smoke_threshold)
            #= brm-docs-example.jl:25 =#
            st
        end
end
stan
functions {
matrix hcat(vector x, vector y, vector z) {
    return hcat(hcat(x, y), z);
}
matrix hcat(
    matrix x,
    vector y
) {
    int m = dims(x)[1];
    int n = dims(x)[2];
    if (dims(y)[1] != m) reject("hcat: dim mismatch — `y` dim 1 (= ", dims(y)[1], ") does not match `m` (= ", m, "), inferred from `x` dim 1. `m` sizes: `x` dim 1 (= ", dims(x)[1], "), `y` dim 1 (= ", dims(y)[1], ").");
    return append_col(x, y);
}
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);
}
matrix hcat(vector x) {
    int n = dims(x)[1];
    return to_matrix(x, n, 1);
}
int ragged_end(array[] int ends, int i) {
    return ends[i];
}
int ragged_start(
    array[] int ends,
    int i
) {
    if((i == 1)) {
        return 1;
    } else {
        return (1 + ends[(i - 1)]);
    }
}
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 normal_lpdfs(
    vector obs,
    vector loc,
    real scale
) {
    return jbroadcasted_normal_lpdfs(obs, loc, scale);
}
vector jbroadcasted_normal_lpdfs(
    vector x1,
    vector x2,
    real x3
) {
    int n = dims(x1)[1];
    vector[n] rv;
    for(i in 1:n) {
        rv[i] = normal_lpdfs(broadcasted_getindex(x1, i), broadcasted_getindex(x2, i), x3);
    }
    return rv;
}
real normal_lpdfs(
    real args1,
    real args2,
    real args3
) {
    return normal_lpdf(args1 | args2, args3);
}
real broadcasted_getindex(vector x, int i) {
    return x[i];
}
vector normal_vector_rng(
    int anontok__1,
    vector a,
    real b
) {
    int n = anontok__1;
    if((n == 0)) {
        vector[n] rv;
        return rv;
    } else {
        return to_vector(normal_rng(a, b));
    }
}
vector bernoulli_logit_lpmfs(
    array[] int obs,
    vector args1
) {
    return jbroadcasted_bernoulli_logit_lpmfs(obs, args1);
}
vector jbroadcasted_bernoulli_logit_lpmfs(
    array[] int x1,
    vector x2
) {
    int n = dims(x1)[1];
    vector[n] rv;
    for(i in 1:n) {
        rv[i] = bernoulli_logit_lpmfs(broadcasted_getindex(x1, i), broadcasted_getindex(x2, i));
    }
    return rv;
}
real bernoulli_logit_lpmfs(
    int args1,
    real args2
) {
    return bernoulli_logit_lpmf(args1 | args2);
}
int broadcasted_getindex(array[] int x, int i) {
    return x[i];
}
array[] int bernoulli_logit_int_rng(
    int anontok__1,
    vector p
) {
    int n = anontok__1;
    return bernoulli_logit_rng(p);
}
vector ema_path(
    vector dt,
    vector wl,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q0,
    real qw,
    real diffm,
    real diff21,
    real s0,
    real m0,
    vector zs,
    vector zm
) {
    int nt = dims(dt)[1];
    if (dims(wl)[1] != nt) reject("ema_path: dim mismatch — `wl` dim 1 (= ", dims(wl)[1], ") does not match `nt` (= ", nt, "), inferred from `dt` dim 1. `nt` sizes: `dt` dim 1 (= ", dims(dt)[1], "), `wl` dim 1 (= ", dims(wl)[1], "), `zs` dim 1 (= ", dims(zs)[1], "), `zm` dim 1 (= ", dims(zm)[1], ").");
    if (dims(zs)[1] != nt) reject("ema_path: dim mismatch — `zs` dim 1 (= ", dims(zs)[1], ") does not match `nt` (= ", nt, "), inferred from `dt` dim 1. `nt` sizes: `dt` dim 1 (= ", dims(dt)[1], "), `wl` dim 1 (= ", dims(wl)[1], "), `zs` dim 1 (= ", dims(zs)[1], "), `zm` dim 1 (= ", dims(zm)[1], ").");
    if (dims(zm)[1] != nt) reject("ema_path: dim mismatch — `zm` dim 1 (= ", dims(zm)[1], ") does not match `nt` (= ", nt, "), inferred from `dt` dim 1. `nt` sizes: `dt` dim 1 (= ", dims(dt)[1], "), `wl` dim 1 (= ", dims(wl)[1], "), `zs` dim 1 (= ", dims(zs)[1], "), `zm` dim 1 (= ", dims(zm)[1], ").");
    vector[(2 * nt)] out;
    real s = (s0 + (wls * wl[1]));
    real m = m0;
    real corr = tanh(diff21);
    out[1] = s;
    out[(nt + 1)] = m;
    for(t in 2:nt) {
        real drift_s = (((-log1p(exp((b0 + (bm * m))))) * s) + (a12 * m));
        real drift_m = ((a21 * s) + (a22 * m) + cm);
        real gs = exp((q0 + (qw * wl[t])));
        real zc = ((corr * zs[t]) + (sqrt((1 - (corr * corr))) * zm[t]));
        s = (s + (drift_s * dt[t]) + (gs * sqrt(dt[t]) * zs[t]) + (wls * wl[t]));
        m = (m + (drift_m * dt[t]) + (diffm * sqrt(dt[t]) * zc));
        out[t] = s;
        out[(nt + t)] = m;
    }
    return out;
}
}
data {
    int n_terms_p_subject;
    int n_subject;
    int treatment_n;
    int age_n;
    vector[age_n] age;
    vector[treatment_n] treatment;
    int subject_idx_n;
    array[subject_idx_n] int subject_idx;
    int kernel_nsub_stress;
    int dt_grid_ends_n;
    int dt_grid_mem_n;
    tuple(vector[dt_grid_mem_n], array[dt_grid_ends_n] int) dt_grid;
    int stressReport_mem_n;
    int stressReport_ends_n;
    tuple(vector[stressReport_mem_n], array[stressReport_ends_n] int) stressReport;
    int moodReport_mem_n;
    int moodReport_ends_n;
    tuple(vector[moodReport_mem_n], array[moodReport_ends_n] int) moodReport;
    int smoked_mem_n;
    int smoked_ends_n;
    tuple(array[smoked_mem_n] int, array[smoked_ends_n] int) smoked;
    int workload_ends_n;
    int workload_mem_n;
    tuple(vector[workload_mem_n], array[workload_ends_n] int) workload;
}
transformed data {
    matrix[treatment_n, (2 + 1)] X_b0 = hcat(rep_vector(1.0, num_elements(age)), age, treatment);
    int pop_b0_n_covariates = (2 + 1);
    matrix[num_elements(subject_idx), 1] X_q0 = hcat(rep_vector(1.0, num_elements(subject_idx)));
    int pop_q0_n_covariates = 1;
    matrix[treatment_n, 2] X_cm = hcat(rep_vector(1.0, num_elements(treatment)), treatment);
    int pop_cm_n_covariates = 2;
    matrix[num_elements(subject_idx), 1] X_wls = hcat(rep_vector(1.0, num_elements(subject_idx)));
    int pop_wls_n_covariates = 1;
    matrix[num_elements(subject_idx), 1] X_s0 = hcat(rep_vector(1.0, num_elements(subject_idx)));
    int pop_s0_n_covariates = 1;
    matrix[num_elements(subject_idx), 1] X_m0 = hcat(rep_vector(1.0, num_elements(subject_idx)));
    int pop_m0_n_covariates = 1;
    array[kernel_nsub_stress] int stress_zs__pl_len_1;
    array[kernel_nsub_stress] int stress_path__pl_len_1;
    array[kernel_nsub_stress] int stress_zm__pl_len_1;
    array[kernel_nsub_stress] int stress__pl_len_1;
    array[kernel_nsub_stress] int stress_st__pl_len_1;
    array[kernel_nsub_stress] int stress_mo__pl_len_1;
    for(plate_i__pl_1 in 1:kernel_nsub_stress) {
        stress_zs__pl_len_1[plate_i__pl_1] = (1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1)));
        stress_path__pl_len_1[plate_i__pl_1] = (2 * (1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1))));
        stress_zm__pl_len_1[plate_i__pl_1] = (1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1)));
        stress__pl_len_1[plate_i__pl_1] = (1 + ((1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1))) - 1));
        stress_st__pl_len_1[plate_i__pl_1] = (1 + ((1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1))) - 1));
        stress_mo__pl_len_1[plate_i__pl_1] = (
            1 +
            (
                (2 * (1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1)))) -
                ((1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1))) + 1)
            )
        );
    }
    array[kernel_nsub_stress] int stress_zs__pl_end_1 = cumulative_sum(stress_zs__pl_len_1);
    array[kernel_nsub_stress] int stress_path__pl_end_1 = cumulative_sum(stress_path__pl_len_1);
    array[kernel_nsub_stress] int stress_zm__pl_end_1 = cumulative_sum(stress_zm__pl_len_1);
    array[kernel_nsub_stress] int stress__pl_end_1 = cumulative_sum(stress__pl_len_1);
    array[kernel_nsub_stress] int stress_st__pl_end_1 = cumulative_sum(stress_st__pl_len_1);
    array[kernel_nsub_stress] int stress_mo__pl_end_1 = cumulative_sum(stress_mo__pl_len_1);
}
parameters {
    cholesky_factor_corr[n_terms_p_subject] b_p_subject_L;
    vector<lower=0.0>[n_terms_p_subject] b_p_subject_tau;
    vector[(n_terms_p_subject * n_subject)] b_p_subject_z_flat;
    real<lower=0.0> sigma_s;
    real<lower=0.0> sigma_m;
    real bm;
    real a12;
    real a21;
    real a22;
    real qw;
    real<lower=0.0> diffm;
    real diff21;
    real l31;
    real smoke_threshold;
    real mm_s;
    real mm_m;
    vector[pop_b0_n_covariates] pop_b0_beta_pop;
    vector[pop_q0_n_covariates] pop_q0_beta_pop;
    vector[pop_cm_n_covariates] pop_cm_beta_pop;
    vector[pop_wls_n_covariates] pop_wls_beta_pop;
    vector[pop_s0_n_covariates] pop_s0_beta_pop;
    real r_s0_subject_log_scale;
    vector[n_subject] r_s0_subject_xi;
    vector[pop_m0_n_covariates] pop_m0_beta_pop;
    real r_m0_subject_log_scale;
    vector[n_subject] r_m0_subject_xi;
    vector[sum(stress_zs__pl_len_1)] stress_zs__pl_mem_1;
    vector[sum(stress_zm__pl_len_1)] stress_zm__pl_mem_1;
}
transformed parameters {
    matrix[n_terms_p_subject, n_subject] b_p_subject_z = to_matrix(b_p_subject_z_flat, n_terms_p_subject, n_subject);
    matrix[n_subject, n_terms_p_subject] b_p_subject = ((diag_pre_multiply(b_p_subject_tau, b_p_subject_L) * b_p_subject_z)');
    vector[treatment_n] pop_b0 = (X_b0 * pop_b0_beta_pop);
    vector[subject_idx_n] r_b0_p_subject = b_p_subject[subject_idx, 1];
    vector[treatment_n] b0 = (pop_b0 + r_b0_p_subject);
    vector[num_elements(subject_idx)] pop_q0 = (X_q0 * pop_q0_beta_pop);
    vector[subject_idx_n] r_q0_p_subject = b_p_subject[subject_idx, 2];
    vector[num_elements(subject_idx)] q0 = (pop_q0 + r_q0_p_subject);
    vector[treatment_n] pop_cm = (X_cm * pop_cm_beta_pop);
    vector[subject_idx_n] r_cm_p_subject = b_p_subject[subject_idx, 3];
    vector[treatment_n] cm = (pop_cm + r_cm_p_subject);
    vector[num_elements(subject_idx)] pop_wls = (X_wls * pop_wls_beta_pop);
    vector[subject_idx_n] r_wls_p_subject = b_p_subject[subject_idx, 4];
    vector[num_elements(subject_idx)] wls = (pop_wls + r_wls_p_subject);
    vector[num_elements(subject_idx)] pop_s0 = (X_s0 * pop_s0_beta_pop);
    vector[subject_idx_n] r_s0_subject = (exp(r_s0_subject_log_scale) * r_s0_subject_xi[subject_idx]);
    vector[num_elements(subject_idx)] s0 = (pop_s0 + r_s0_subject);
    vector[num_elements(subject_idx)] pop_m0 = (X_m0 * pop_m0_beta_pop);
    vector[subject_idx_n] r_m0_subject = (exp(r_m0_subject_log_scale) * r_m0_subject_xi[subject_idx]);
    vector[num_elements(subject_idx)] m0 = (pop_m0 + r_m0_subject);
    vector[sum(stress_path__pl_len_1)] stress_path__pl_mem_1;
    vector[sum(stress_st__pl_len_1)] stress_st__pl_mem_1;
    vector[sum(stress_mo__pl_len_1)] stress_mo__pl_mem_1;
    for(plate_i__pl_1 in 1:kernel_nsub_stress) {
        stress_path__pl_mem_1[
            ragged_start(stress_path__pl_end_1, plate_i__pl_1):ragged_end(stress_path__pl_end_1, plate_i__pl_1)
        ] = ema_path(
            dt_grid.1[ragged_start(dt_grid.2, plate_i__pl_1):ragged_end(dt_grid.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            b0[plate_i__pl_1],
            bm,
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls[plate_i__pl_1],
            q0[plate_i__pl_1],
            qw,
            diffm,
            diff21,
            s0[plate_i__pl_1],
            m0[plate_i__pl_1],
            stress_zs__pl_mem_1[
                ragged_start(stress_zs__pl_end_1, plate_i__pl_1):ragged_end(stress_zs__pl_end_1, plate_i__pl_1)
            ],
            stress_zm__pl_mem_1[
                ragged_start(stress_zm__pl_end_1, plate_i__pl_1):ragged_end(stress_zm__pl_end_1, plate_i__pl_1)
            ]
        );
        stress_st__pl_mem_1[
            ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
        ] = stress_path__pl_mem_1[
            ragged_start(stress_path__pl_end_1, plate_i__pl_1):ragged_end(stress_path__pl_end_1, plate_i__pl_1)
        ][
            1:(1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1)))
        ];
        stress_mo__pl_mem_1[
            ragged_start(stress_mo__pl_end_1, plate_i__pl_1):ragged_end(stress_mo__pl_end_1, plate_i__pl_1)
        ] = stress_path__pl_mem_1[
            ragged_start(stress_path__pl_end_1, plate_i__pl_1):ragged_end(stress_path__pl_end_1, plate_i__pl_1)
        ][
            ((1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1))) + 1):(2 * (1 + (ragged_end(dt_grid.2, plate_i__pl_1) - ragged_start(dt_grid.2, plate_i__pl_1))))
        ];
    }
}
model {
    b_p_subject_L ~ lkj_corr_cholesky(1.0);
    b_p_subject_tau ~ std_normal();
    b_p_subject_z_flat ~ std_normal();
    sigma_s ~ exponential((1.0 ./ 1));
    sigma_m ~ exponential((1.0 ./ 1));
    bm ~ normal(0, 0.5);
    a12 ~ normal(0, 0.5);
    a21 ~ normal(0, 0.5);
    a22 ~ normal(-0.5, 0.3);
    qw ~ normal(0, 0.5);
    diffm ~ exponential((1.0 ./ 1));
    diff21 ~ normal(0, 0.5);
    l31 ~ normal(0, 1);
    smoke_threshold ~ normal(0, 1);
    mm_s ~ normal(0, 0.5);
    mm_m ~ normal(0, 0.5);
    pop_b0_beta_pop ~ std_normal();
    pop_q0_beta_pop ~ std_normal();
    pop_cm_beta_pop ~ std_normal();
    pop_wls_beta_pop ~ std_normal();
    pop_s0_beta_pop ~ std_normal();
    r_s0_subject_log_scale ~ std_normal();
    r_s0_subject_xi ~ std_normal();
    pop_m0_beta_pop ~ std_normal();
    r_m0_subject_log_scale ~ std_normal();
    r_m0_subject_xi ~ std_normal();
    for(plate_i__pl_1 in 1:kernel_nsub_stress) {
        stress_zs__pl_mem_1[
            ragged_start(stress_zs__pl_end_1, plate_i__pl_1):ragged_end(stress_zs__pl_end_1, plate_i__pl_1)
        ] ~ std_normal();
        stress_zm__pl_mem_1[
            ragged_start(stress_zm__pl_end_1, plate_i__pl_1):ragged_end(stress_zm__pl_end_1, plate_i__pl_1)
        ] ~ std_normal();
        stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] ~ normal(
            (
                mm_s +
                stress_st__pl_mem_1[
                    ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
                ]
            ),
            sigma_s
        );
        moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)] ~ normal(
            (
                mm_m +
                stress_mo__pl_mem_1[
                    ragged_start(stress_mo__pl_end_1, plate_i__pl_1):ragged_end(stress_mo__pl_end_1, plate_i__pl_1)
                ]
            ),
            sigma_m
        );
        smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)] ~ bernoulli_logit(
            (
                (
                    l31 .*
                    stress_st__pl_mem_1[
                        ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
                    ]
                ) +
                smoke_threshold
            )
        );
    }
}
generated quantities {
    vector[sum(stress__pl_len_1)] stress__pl_mem_1;
    vector[num_elements(stressReport.1)] stressReport_gen;
    vector[num_elements(stressReport.2)] stressReport_likelihood;
    vector[num_elements(moodReport.1)] moodReport_gen;
    vector[num_elements(moodReport.2)] moodReport_likelihood;
    array[num_elements(smoked.1)] int smoked_gen;
    vector[num_elements(smoked.2)] smoked_likelihood;
    for(plate_i__pl_1 in 1:kernel_nsub_stress) {
        stressReport_gen[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] = normal_vector_rng(
            (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1))),
            (
                mm_s +
                stress_st__pl_mem_1[
                    ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
                ]
            ),
            sigma_s
        );
        stressReport_likelihood[plate_i__pl_1] = normal_lpdf(stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] | 
            (
                mm_s +
                stress_st__pl_mem_1[
                    ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
                ]
            ),
            sigma_s
        );
        moodReport_gen[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)] = normal_vector_rng(
            (1 + (ragged_end(moodReport.2, plate_i__pl_1) - ragged_start(moodReport.2, plate_i__pl_1))),
            (
                mm_m +
                stress_mo__pl_mem_1[
                    ragged_start(stress_mo__pl_end_1, plate_i__pl_1):ragged_end(stress_mo__pl_end_1, plate_i__pl_1)
                ]
            ),
            sigma_m
        );
        moodReport_likelihood[plate_i__pl_1] = normal_lpdf(moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)] | 
            (
                mm_m +
                stress_mo__pl_mem_1[
                    ragged_start(stress_mo__pl_end_1, plate_i__pl_1):ragged_end(stress_mo__pl_end_1, plate_i__pl_1)
                ]
            ),
            sigma_m
        );
        smoked_gen[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)] = bernoulli_logit_int_rng(
            (1 + (ragged_end(smoked.2, plate_i__pl_1) - ragged_start(smoked.2, plate_i__pl_1))),
            (
                (
                    l31 .*
                    stress_st__pl_mem_1[
                        ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
                    ]
                ) +
                smoke_threshold
            )
        );
        smoked_likelihood[plate_i__pl_1] = bernoulli_logit_lpmf(smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)] | 
            (
                (
                    l31 .*
                    stress_st__pl_mem_1[
                        ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
                    ]
                ) +
                smoke_threshold
            )
        );
        stress__pl_mem_1[
            ragged_start(stress__pl_end_1, plate_i__pl_1):ragged_end(stress__pl_end_1, plate_i__pl_1)
        ] = stress_st__pl_mem_1[
            ragged_start(stress_st__pl_end_1, plate_i__pl_1):ragged_end(stress_st__pl_end_1, plate_i__pl_1)
        ];
    }
}
julia
Turing unsupported for this BRM example

Turing backend: direct execution requires at least one observed likelihood

On the 6-subject × 10-occasion fixture this model has 190 parameters, most of them innovations: its dimension grows with subjects × occasions × processes, and the innovations meet their scale parameters in the funnel geometry familiar from every non-marginalized hierarchical model. It is the most flexible form — any observation family, any nonlinearity — and the most expensive one. (The Turing pane is retained on this page even though kernel cells are outside the current Turing executor; its construction error documents that backend boundary.)

2. Integrating the states out exactly: a Kalman filter in the cell ​

If the dynamics are linear and everything is Gaussian, the latent path can be integrated out in closed form. The filter becomes a custom likelihood family: a StanBlocks @lpxf triad — _lpdf (the log-likelihood), _lpdfs (its pointwise terms, for LOO and hold-out) and _rng (posterior-predictive draws) — that the cell calls once per subject as ys ~ kalman2(...).

julia
StanBlocks.@deffun begin
    l2pi()::real = 1.8378770664093453          # log(2 pi)
    @lhs @lpxf kalman2_lpdf(ys::vector[T], ym::vector[T], workload::vector[T], dt::vector[T],
            a11::real, a12::real, a21::real, a22::real, cm::real, wls::real,
            q1::real, q2::real, r1::real, r2::real,
            ms0::real, mm0::real, P0::real)::real = begin
        ms=ms0; mm=mm0; p11=P0; p12=0.0; p22=P0; ll=0.0
        for t in 1:T
            if t>1
                d=dt[t]
                nms=ms+(a11*ms+a12*mm)*d; nmm=mm+(a21*ms+a22*mm+cm)*d
                f11=1+a11*d; f12=a12*d; f21=a21*d; f22=1+a22*d
                qs=q1*q1*d; qm=q2*q2*d
                fp11=f11*p11+f12*p12; fp12=f11*p12+f12*p22; fp21=f21*p11+f22*p12; fp22=f21*p12+f22*p22
                np11=fp11*f11+fp12*f12+qs; np12=fp11*f21+fp12*f22; np22=fp21*f21+fp22*f22+qm
                ms=nms; mm=nmm; p11=np11; p12=np12; p22=np22
            end
            ms=ms+wls*workload[t]                    # workload acts as an impulse at the observation
            v1=ys[t]-ms; v2=ym[t]-mm; s11=p11+r1*r1; s12=p12; s22=p22+r2*r2   # r1, r2: measurement sds
            det=s11*s22-s12*s12; si11=s22/det; si12=-s12/det; si22=s11/det
            quad=v1*(si11*v1+si12*v2)+v2*(si12*v1+si22*v2)
            ll=ll-0.5*(2*l2pi()+log(det)+quad)
            k11=p11*si11+p12*si12; k12=p11*si12+p12*si22; k21=p12*si11+p22*si12; k22=p12*si12+p22*si22
            ms=ms+k11*v1+k12*v2; mm=mm+k21*v1+k22*v2
            g11=(1-k11)*p11-k12*p12; g12=(1-k11)*p12-k12*p22; g21=-k21*p11+(1-k22)*p12; g22=-k21*p12+(1-k22)*p22
            p11=g11; p12=0.5*(g12+g21); p22=g22
        end
        ll
    end
    kalman2_lpdfs(ys::vector[T], ym::vector[T], workload::vector[T], dt::vector[T],
            a11::real, a12::real, a21::real, a22::real, cm::real, wls::real,
            q1::real, q2::real, r1::real, r2::real,
            ms0::real, mm0::real, P0::real)::vector[T] = begin
        out::vector[T]; ms=ms0; mm=mm0; p11=P0; p12=0.0; p22=P0
        for t in 1:T
            if t>1
                d=dt[t]
                nms=ms+(a11*ms+a12*mm)*d; nmm=mm+(a21*ms+a22*mm+cm)*d
                f11=1+a11*d; f12=a12*d; f21=a21*d; f22=1+a22*d; qs=q1*q1*d; qm=q2*q2*d
                fp11=f11*p11+f12*p12; fp12=f11*p12+f12*p22; fp21=f21*p11+f22*p12; fp22=f21*p12+f22*p22
                np11=fp11*f11+fp12*f12+qs; np12=fp11*f21+fp12*f22; np22=fp21*f21+fp22*f22+qm
                ms=nms; mm=nmm; p11=np11; p12=np12; p22=np22
            end
            ms=ms+wls*workload[t]                    # workload acts as an impulse at the observation
            v1=ys[t]-ms; v2=ym[t]-mm; s11=p11+r1*r1; s12=p12; s22=p22+r2*r2   # r1, r2: measurement sds
            det=s11*s22-s12*s12; si11=s22/det; si12=-s12/det; si22=s11/det
            quad=v1*(si11*v1+si12*v2)+v2*(si12*v1+si22*v2); out[t]=-0.5*(2*l2pi()+log(det)+quad)
            k11=p11*si11+p12*si12; k12=p11*si12+p12*si22; k21=p12*si11+p22*si12; k22=p12*si12+p22*si22
            ms=ms+k11*v1+k12*v2; mm=mm+k21*v1+k22*v2
            g11=(1-k11)*p11-k12*p12; g12=(1-k11)*p12-k12*p22; g21=-k21*p11+(1-k22)*p12; g22=-k21*p12+(1-k22)*p22
            p11=g11; p12=0.5*(g12+g21); p22=g22
        end
        out
    end
    kalman2_rng(vector[T], ym::vector[T], workload::vector[T], dt::vector[T],
            a11::real, a12::real, a21::real, a22::real, cm::real, wls::real,
            q1::real, q2::real, r1::real, r2::real,
            ms0::real, mm0::real, P0::real)::vector[T] = begin
        out::vector[T]; s=normal_rng(ms0,sqrt(P0)); m=normal_rng(mm0,sqrt(P0))
        for t in 1:T
            if t>1
                d=dt[t]
                s=s+(a11*s+a12*m)*d+q1*sqrt(d)*normal_rng(0.,1.)
                m=m+(a21*s+a22*m+cm)*d+q2*sqrt(d)*normal_rng(0.,1.)
            end
            s=s+wls*workload[t]
            out[t]=normal_rng(s,r1)
        end
        out
    end
end
brm-comparison
Linear-Gaussian model, states Kalman-marginalized per subject
julia
function ema_kernel_kalman_model(data = ema_kalman_fixture())
    @brm data begin
        # shared drift / coupling / noise parameters
        a12 ~ Normal(0, 0.5)                   # mood -> stress
        a21 ~ Normal(0, 0.5)                   # stress -> mood
        a22 ~ Normal(-0.5, 0.3)                # mood self-decay
        wls ~ Normal(0, 0.5)                   # workload -> stress
        q1  ~ Exponential(1.0)                 # stress process-noise sd
        q2  ~ Exponential(1.0)                 # mood process-noise sd
        r1  ~ Exponential(1.0)                 # stressReport measurement sd
        r2  ~ Exponential(1.0)                 # moodReport measurement sd

        # subject-level parameters: covariates + random effects on the formula surface
        a11 ~ 1 + age       + (1 | subject)    # stress self-decay
        cm  ~ 1 + treatment + (1 | subject)    # mood intercept
        s0  ~ 1             + (1 | subject)    # initial stress mean
        m0  ~ 1             + (1 | subject)    # initial mood mean

        # per subject: integrate out THIS subject's entire latent path, exactly
        pred ~ kernel(dt, workload, stressReport, moodReport,
                      a11, cm, s0, m0) do dti, wli, ys, ym, la11, lcm, ls0, lm0
            ys ~ kalman2(ym, wli, dti, la11, a12, a21, a22, lcm, wls, q1, q2, r1, r2, ls0, lm0, 10.0)
            ys
        end
    end
end
julia
BRMI:
  a12 ~ Normal(0, 0.5)
  a21 ~ Normal(0, 0.5)
  a22 ~ Normal(-0.5, 0.3)
  wls ~ Normal(0, 0.5)
  q1 ~ Exponential(1.0)
  q2 ~ Exponential(1.0)
  r1 ~ Exponential(1.0)
  r2 ~ Exponential(1.0)
  age: data (eltype=Float64, n=5)
  subject: data (eltype=String, n=5)
  a11 ~ 1 + age + (1 | subject)
  treatment: data (eltype=Float64, n=5)
  cm ~ 1 + treatment + (1 | subject)
  s0 ~ 1 + (1 | subject)
  m0 ~ 1 + (1 | subject)
  dt: data (eltype=Vector{Float64}, n=5)
  workload: data (eltype=Vector{Float64}, n=5)
  stressReport: data (eltype=Vector{Float64}, n=5)
  moodReport: data (eltype=Vector{Float64}, n=5)
  pred ~ kernel((dti, wli, ys, ym, la11, lcm, ls0, lm0)->begin
        #= brm-docs-example.jl:22 =#
        ys ~ kalman2(ym, wli, dti, la11, a12, a21, a22, lcm, wls, q1, q2, r1, r2, ls0, lm0, 10.0)
        #= brm-docs-example.jl:23 =#
        ys
    end, dt, workload, stressReport, moodReport, a11, cm, s0, m0)
julia
SBBRMI with data keys = [:age, :dt, :kernel_nsub_pred, :moodReport, :stressReport, :total_A_a11, :total_A_cm, :total_A_m0, :total_A_s0, :total_group_a11, :total_group_cm, :total_group_m0, :total_group_s0, :total_location_a11, :total_location_cm, :total_location_m0, :total_location_s0, :total_ng_a11, :total_ng_cm, :total_ng_m0, :total_ng_s0, :total_nk_a11, :total_nk_cm, :total_nk_m0, :total_nk_s0, :total_np_a11, :total_np_cm, :total_np_m0, :total_np_s0, :total_precision_a11, :total_precision_cm, :total_precision_m0, :total_precision_s0, :treatment, :workload]
configured submodels:
_brm_total_scales_configured_1 = Base.merge(BayesianRegressionModels._brm_total_scales, quote
            tau ~ (ValueFamily(brm_vector_prior_faeb6f6956d0662a))(0.0, 1.0; n = 1)
        end)
emitted @slic body:
begin
    a12 ~ normal(0, 0.5)
    a21 ~ normal(0, 0.5)
    a22 ~ normal(-0.5, 0.3)
    wls ~ normal(0, 0.5)
    q1 ~ exponential(1.0 ./ 1.0)
    q2 ~ exponential(1.0 ./ 1.0)
    r1 ~ exponential(1.0 ./ 1.0)
    r2 ~ exponential(1.0 ./ 1.0)
    total_scale_a11 ~ _brm_total_scales_configured_1(; n = total_nk_a11)
    total_a11::matrix[total_ng_a11, total_nk_a11] ~ brm_total(total_scale_a11, total_A_a11, total_location_a11, total_precision_a11)
    population_a11 = brm_total_recover_rng(total_a11, total_scale_a11, total_A_a11, total_location_a11, total_precision_a11)
    deviation_a11 = brm_total_deviations(total_a11, total_A_a11 * population_a11)
    total_Z_a11 = hcat(rep_vector(1.0, num_elements(total_group_a11)))
    X_a11 = hcat(age)
    pop_a11 ~ popefs(; X = X_a11)
    a11 = rows_dot_product(total_a11[total_group_a11, :], total_Z_a11) + pop_a11
    total_scale_cm ~ _brm_total_scales_configured_1(; n = total_nk_cm)
    total_cm::matrix[total_ng_cm, total_nk_cm] ~ brm_total(total_scale_cm, total_A_cm, total_location_cm, total_precision_cm)
    population_cm = brm_total_recover_rng(total_cm, total_scale_cm, total_A_cm, total_location_cm, total_precision_cm)
    deviation_cm = brm_total_deviations(total_cm, total_A_cm * population_cm)
    total_Z_cm = hcat(rep_vector(1.0, num_elements(total_group_cm)))
    X_cm = hcat(treatment)
    pop_cm ~ popefs(; X = X_cm)
    cm = rows_dot_product(total_cm[total_group_cm, :], total_Z_cm) + pop_cm
    total_scale_s0 ~ _brm_total_scales_configured_1(; n = total_nk_s0)
    total_s0::matrix[total_ng_s0, total_nk_s0] ~ brm_total(total_scale_s0, total_A_s0, total_location_s0, total_precision_s0)
    population_s0 = brm_total_recover_rng(total_s0, total_scale_s0, total_A_s0, total_location_s0, total_precision_s0)
    deviation_s0 = brm_total_deviations(total_s0, total_A_s0 * population_s0)
    total_Z_s0 = hcat(rep_vector(1.0, num_elements(total_group_s0)))
    s0 = rows_dot_product(total_s0[total_group_s0, :], total_Z_s0)
    total_scale_m0 ~ _brm_total_scales_configured_1(; n = total_nk_m0)
    total_m0::matrix[total_ng_m0, total_nk_m0] ~ brm_total(total_scale_m0, total_A_m0, total_location_m0, total_precision_m0)
    population_m0 = brm_total_recover_rng(total_m0, total_scale_m0, total_A_m0, total_location_m0, total_precision_m0)
    deviation_m0 = brm_total_deviations(total_m0, total_A_m0 * population_m0)
    total_Z_m0 = hcat(rep_vector(1.0, num_elements(total_group_m0)))
    m0 = rows_dot_product(total_m0[total_group_m0, :], total_Z_m0)
    pred ~ plate(dt, workload, stressReport, moodReport, a11, cm, s0, m0; outer = (kernel_nsub_pred,)) do dti, wli, ys, ym, la11, lcm, ls0, lm0
            #= brm-docs-example.jl:22 =#
            ys ~ kalman2(ym, wli, dti, la11, a12, a21, a22, lcm, wls, q1, q2, r1, r2, ls0, lm0, 10.0)
            #= brm-docs-example.jl:23 =#
            ys
        end
end
stan
functions {
// value UDF brm_vector_prior_faeb6f6956d0662a_lpdf
real brm_vector_prior_faeb6f6956d0662a_lpdf(
    vector x,
    real arg_1,
    real arg_2
) {
    if((x[1] < 0.0)) {
        return negative_infinity();
    }
    return lognormal_lpdf(x[1] | arg_1, arg_2);
}
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) {
    int n = dims(x)[1];
    return to_matrix(x, n, 1);
}
int ragged_end(array[] int ends, int i) {
    return ends[i];
}
int ragged_start(
    array[] int ends,
    int i
) {
    if((i == 1)) {
        return 1;
    } else {
        return (1 + ends[(i - 1)]);
    }
}
real kalman2_lpdf(
    vector ys,
    vector ym,
    vector workload,
    vector dt,
    real a11,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q1,
    real q2,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real P0
) {
    int T = dims(ys)[1];
    if (dims(ym)[1] != T) reject("kalman2_lpdf: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(workload)[1] != T) reject("kalman2_lpdf: dim mismatch — `workload` dim 1 (= ", dims(workload)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("kalman2_lpdf: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    real ms = ms0;
    real mm = mm0;
    real p11 = P0;
    real p12 = 0.0;
    real p22 = P0;
    real ll = 0.0;
    for(t in 1:T) {
        if((t > 1)) {
            real d = dt[t];
            real nms = (ms + (((a11 * ms) + (a12 * mm)) * d));
            real nmm = (mm + (((a21 * ms) + (a22 * mm) + cm) * d));
            real f11 = (1 + (a11 * d));
            real f12 = (a12 * d);
            real f21 = (a21 * d);
            real f22 = (1 + (a22 * d));
            real qs = (q1 * q1 * d);
            real qm = (q2 * q2 * d);
            real fp11 = ((f11 * p11) + (f12 * p12));
            real fp12 = ((f11 * p12) + (f12 * p22));
            real fp21 = ((f21 * p11) + (f22 * p12));
            real fp22 = ((f21 * p12) + (f22 * p22));
            real np11 = ((fp11 * f11) + (fp12 * f12) + qs);
            real np12 = ((fp11 * f21) + (fp12 * f22));
            real np22 = ((fp21 * f21) + (fp22 * f22) + qm);
            ms = nms;
            mm = nmm;
            p11 = np11;
            p12 = np12;
            p22 = np22;
        }
        ms = (ms + (wls * workload[t]));
        real v1 = (ys[t] - ms);
        real v2 = (ym[t] - mm);
        real s11 = (p11 + (r1 * r1));
        real s12 = p12;
        real s22 = (p22 + (r2 * r2));
        real det = ((s11 * s22) - (s12 * s12));
        real si11 = (s22 / det);
        real si12 = ((-s12) / det);
        real si22 = (s11 / det);
        real quad = ((v1 * ((si11 * v1) + (si12 * v2))) + (v2 * ((si12 * v1) + (si22 * v2))));
        ll = (ll - (0.5 * ((2 * l2pi()) + log(det) + quad)));
        real k11 = ((p11 * si11) + (p12 * si12));
        real k12 = ((p11 * si12) + (p12 * si22));
        real k21 = ((p12 * si11) + (p22 * si12));
        real k22 = ((p12 * si12) + (p22 * si22));
        ms = (ms + (k11 * v1) + (k12 * v2));
        mm = (mm + (k21 * v1) + (k22 * v2));
        real g11 = (((1 - k11) * p11) - (k12 * p12));
        real g12 = (((1 - k11) * p12) - (k12 * p22));
        real g21 = (((-k21) * p11) + ((1 - k22) * p12));
        real g22 = (((-k21) * p12) + ((1 - k22) * p22));
        p11 = g11;
        p12 = (0.5 * (g12 + g21));
        p22 = g22;
    }
    return ll;
}
real l2pi() {
    return 1.8378770664093453;
}
vector kalman2_lpdfs(
    vector ys,
    vector ym,
    vector workload,
    vector dt,
    real a11,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q1,
    real q2,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real P0
) {
    int T = dims(ys)[1];
    if (dims(ym)[1] != T) reject("kalman2_lpdfs: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(workload)[1] != T) reject("kalman2_lpdfs: dim mismatch — `workload` dim 1 (= ", dims(workload)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("kalman2_lpdfs: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    vector[T] out;
    real ms = ms0;
    real mm = mm0;
    real p11 = P0;
    real p12 = 0.0;
    real p22 = P0;
    for(t in 1:T) {
        if((t > 1)) {
            real d = dt[t];
            real nms = (ms + (((a11 * ms) + (a12 * mm)) * d));
            real nmm = (mm + (((a21 * ms) + (a22 * mm) + cm) * d));
            real f11 = (1 + (a11 * d));
            real f12 = (a12 * d);
            real f21 = (a21 * d);
            real f22 = (1 + (a22 * d));
            real qs = (q1 * q1 * d);
            real qm = (q2 * q2 * d);
            real fp11 = ((f11 * p11) + (f12 * p12));
            real fp12 = ((f11 * p12) + (f12 * p22));
            real fp21 = ((f21 * p11) + (f22 * p12));
            real fp22 = ((f21 * p12) + (f22 * p22));
            real np11 = ((fp11 * f11) + (fp12 * f12) + qs);
            real np12 = ((fp11 * f21) + (fp12 * f22));
            real np22 = ((fp21 * f21) + (fp22 * f22) + qm);
            ms = nms;
            mm = nmm;
            p11 = np11;
            p12 = np12;
            p22 = np22;
        }
        ms = (ms + (wls * workload[t]));
        real v1 = (ys[t] - ms);
        real v2 = (ym[t] - mm);
        real s11 = (p11 + (r1 * r1));
        real s12 = p12;
        real s22 = (p22 + (r2 * r2));
        real det = ((s11 * s22) - (s12 * s12));
        real si11 = (s22 / det);
        real si12 = ((-s12) / det);
        real si22 = (s11 / det);
        real quad = ((v1 * ((si11 * v1) + (si12 * v2))) + (v2 * ((si12 * v1) + (si22 * v2))));
        out[t] = (-0.5 * ((2 * l2pi()) + log(det) + quad));
        real k11 = ((p11 * si11) + (p12 * si12));
        real k12 = ((p11 * si12) + (p12 * si22));
        real k21 = ((p12 * si11) + (p22 * si12));
        real k22 = ((p12 * si12) + (p22 * si22));
        ms = (ms + (k11 * v1) + (k12 * v2));
        mm = (mm + (k21 * v1) + (k22 * v2));
        real g11 = (((1 - k11) * p11) - (k12 * p12));
        real g12 = (((1 - k11) * p12) - (k12 * p22));
        real g21 = (((-k21) * p11) + ((1 - k22) * p12));
        real g22 = (((-k21) * p12) + ((1 - k22) * p22));
        p11 = g11;
        p12 = (0.5 * (g12 + g21));
        p22 = g22;
    }
    return out;
}
vector kalman2_vector_rng(
    int anontok__1,
    vector ym,
    vector workload,
    vector dt,
    real a11,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q1,
    real q2,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real P0
) {
    int T = anontok__1;
    if (dims(ym)[1] != T) reject("kalman2_rng: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(workload)[1] != T) reject("kalman2_rng: dim mismatch — `workload` dim 1 (= ", dims(workload)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("kalman2_rng: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    vector[T] out;
    real s = normal_rng(ms0, sqrt(P0));
    real m = normal_rng(mm0, sqrt(P0));
    for(t in 1:T) {
        if((t > 1)) {
            real d = dt[t];
            s = (s + (((a11 * s) + (a12 * m)) * d) + (q1 * sqrt(d) * normal_rng(0.0, 1.0)));
            m = (m + (((a21 * s) + (a22 * m) + cm) * d) + (q2 * sqrt(d) * normal_rng(0.0, 1.0)));
        }
        s = (s + (wls * workload[t]));
        out[t] = normal_rng(s, r1);
    }
    return out;
}
}
data {
    int total_ng_a11;
    int total_nk_a11;
    int total_A_a11_m;
    int total_A_a11_n;
    matrix[total_A_a11_m, total_A_a11_n] total_A_a11;
    int total_location_a11_n;
    vector[total_location_a11_n] total_location_a11;
    int total_precision_a11_n;
    vector[total_precision_a11_n] total_precision_a11;
    int total_group_a11_n;
    array[total_group_a11_n] int total_group_a11;
    int age_n;
    vector[age_n] age;
    int total_ng_cm;
    int total_nk_cm;
    int total_A_cm_m;
    int total_A_cm_n;
    matrix[total_A_cm_m, total_A_cm_n] total_A_cm;
    int total_location_cm_n;
    vector[total_location_cm_n] total_location_cm;
    int total_precision_cm_n;
    vector[total_precision_cm_n] total_precision_cm;
    int total_group_cm_n;
    array[total_group_cm_n] int total_group_cm;
    int treatment_n;
    vector[treatment_n] treatment;
    int total_ng_s0;
    int total_nk_s0;
    int total_A_s0_m;
    int total_A_s0_n;
    matrix[total_A_s0_m, total_A_s0_n] total_A_s0;
    int total_location_s0_n;
    vector[total_location_s0_n] total_location_s0;
    int total_precision_s0_n;
    vector[total_precision_s0_n] total_precision_s0;
    int total_group_s0_n;
    array[total_group_s0_n] int total_group_s0;
    int total_ng_m0;
    int total_nk_m0;
    int total_A_m0_m;
    int total_A_m0_n;
    matrix[total_A_m0_m, total_A_m0_n] total_A_m0;
    int total_location_m0_n;
    vector[total_location_m0_n] total_location_m0;
    int total_precision_m0_n;
    vector[total_precision_m0_n] total_precision_m0;
    int total_group_m0_n;
    array[total_group_m0_n] int total_group_m0;
    int kernel_nsub_pred;
    int stressReport_ends_n;
    int stressReport_mem_n;
    tuple(vector[stressReport_mem_n], array[stressReport_ends_n] int) stressReport;
    int dt_ends_n;
    int dt_mem_n;
    tuple(vector[dt_mem_n], array[dt_ends_n] int) dt;
    int moodReport_ends_n;
    int moodReport_mem_n;
    tuple(vector[moodReport_mem_n], array[moodReport_ends_n] int) moodReport;
    int workload_ends_n;
    int workload_mem_n;
    tuple(vector[workload_mem_n], array[workload_ends_n] int) workload;
}
transformed data {
    matrix[num_elements(total_group_a11), 1] total_Z_a11 = hcat(rep_vector(1.0, num_elements(total_group_a11)));
    matrix[age_n, 1] X_a11 = hcat(age);
    int pop_a11_n_covariates = 1;
    matrix[num_elements(total_group_cm), 1] total_Z_cm = hcat(rep_vector(1.0, num_elements(total_group_cm)));
    matrix[treatment_n, 1] X_cm = hcat(treatment);
    int pop_cm_n_covariates = 1;
    matrix[num_elements(total_group_s0), 1] total_Z_s0 = hcat(rep_vector(1.0, num_elements(total_group_s0)));
    matrix[num_elements(total_group_m0), 1] total_Z_m0 = hcat(rep_vector(1.0, num_elements(total_group_m0)));
    array[kernel_nsub_pred] int pred__pl_len_1;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        pred__pl_len_1[plate_i__pl_1] = (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1)));
    }
    array[kernel_nsub_pred] int pred__pl_end_1 = cumulative_sum(pred__pl_len_1);
    vector[sum(pred__pl_len_1)] pred__pl_mem_1;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        pred__pl_mem_1[
            ragged_start(pred__pl_end_1, plate_i__pl_1):ragged_end(pred__pl_end_1, plate_i__pl_1)
        ] = stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ];
    }
}
parameters {
    real a12;
    real a21;
    real a22;
    real wls;
    real<lower=0.0> q1;
    real<lower=0.0> q2;
    real<lower=0.0> r1;
    real<lower=0.0> r2;
    vector<lower=0.0>[1] total_scale_a11_tau;
    matrix[total_ng_a11, total_nk_a11] total_a11;
    vector[pop_a11_n_covariates] pop_a11_beta_pop;
    vector<lower=0.0>[1] total_scale_cm_tau;
    matrix[total_ng_cm, total_nk_cm] total_cm;
    vector[pop_cm_n_covariates] pop_cm_beta_pop;
    vector<lower=0.0>[1] total_scale_s0_tau;
    matrix[total_ng_s0, total_nk_s0] total_s0;
    vector<lower=0.0>[1] total_scale_m0_tau;
    matrix[total_ng_m0, total_nk_m0] total_m0;
}
transformed parameters {
    vector<lower=0.0>[1] total_scale_a11 = total_scale_a11_tau;
    vector[age_n] pop_a11 = (X_a11 * pop_a11_beta_pop);
    vector[num_elements(total_group_a11)] a11 = (rows_dot_product(total_a11[total_group_a11, :], total_Z_a11) + pop_a11);
    vector<lower=0.0>[1] total_scale_cm = total_scale_cm_tau;
    vector[treatment_n] pop_cm = (X_cm * pop_cm_beta_pop);
    vector[num_elements(total_group_cm)] cm = (rows_dot_product(total_cm[total_group_cm, :], total_Z_cm) + pop_cm);
    vector<lower=0.0>[1] total_scale_s0 = total_scale_s0_tau;
    vector[num_elements(total_group_s0)] s0 = rows_dot_product(total_s0[total_group_s0, :], total_Z_s0);
    vector<lower=0.0>[1] total_scale_m0 = total_scale_m0_tau;
    vector[num_elements(total_group_m0)] m0 = rows_dot_product(total_m0[total_group_m0, :], total_Z_m0);
}
model {
    a12 ~ normal(0, 0.5);
    a21 ~ normal(0, 0.5);
    a22 ~ normal(-0.5, 0.3);
    wls ~ normal(0, 0.5);
    q1 ~ exponential((1.0 ./ 1.0));
    q2 ~ exponential((1.0 ./ 1.0));
    r1 ~ exponential((1.0 ./ 1.0));
    r2 ~ exponential((1.0 ./ 1.0));
    total_scale_a11_tau ~ brm_vector_prior_faeb6f6956d0662a(0.0, 1.0);
    total_a11 ~ brm_total(total_scale_a11, total_A_a11, total_location_a11, total_precision_a11);
    pop_a11_beta_pop ~ std_normal();
    total_scale_cm_tau ~ brm_vector_prior_faeb6f6956d0662a(0.0, 1.0);
    total_cm ~ brm_total(total_scale_cm, total_A_cm, total_location_cm, total_precision_cm);
    pop_cm_beta_pop ~ std_normal();
    total_scale_s0_tau ~ brm_vector_prior_faeb6f6956d0662a(0.0, 1.0);
    total_s0 ~ brm_total(total_scale_s0, total_A_s0, total_location_s0, total_precision_s0);
    total_scale_m0_tau ~ brm_vector_prior_faeb6f6956d0662a(0.0, 1.0);
    total_m0 ~ brm_total(total_scale_m0, total_A_m0, total_location_m0, total_precision_m0);
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] ~ kalman2(
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            a11[plate_i__pl_1],
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls,
            q1,
            q2,
            r1,
            r2,
            s0[plate_i__pl_1],
            m0[plate_i__pl_1],
            10.0
        );
    }
}
generated quantities {
    vector[total_precision_a11_n] population_a11 = brm_total_recover_rng(
        total_a11,
        total_scale_a11,
        total_A_a11,
        total_location_a11,
        total_precision_a11
    );
    matrix[total_ng_a11, total_A_a11_m] deviation_a11 = brm_total_deviations(total_a11, (total_A_a11 * population_a11));
    vector[total_precision_cm_n] population_cm = brm_total_recover_rng(total_cm, total_scale_cm, total_A_cm, total_location_cm, total_precision_cm);
    matrix[total_ng_cm, total_A_cm_m] deviation_cm = brm_total_deviations(total_cm, (total_A_cm * population_cm));
    vector[total_precision_s0_n] population_s0 = brm_total_recover_rng(total_s0, total_scale_s0, total_A_s0, total_location_s0, total_precision_s0);
    matrix[total_ng_s0, total_A_s0_m] deviation_s0 = brm_total_deviations(total_s0, (total_A_s0 * population_s0));
    vector[total_precision_m0_n] population_m0 = brm_total_recover_rng(total_m0, total_scale_m0, total_A_m0, total_location_m0, total_precision_m0);
    matrix[total_ng_m0, total_A_m0_m] deviation_m0 = brm_total_deviations(total_m0, (total_A_m0 * population_m0));
    vector[num_elements(stressReport.1)] stressReport_gen;
    vector[num_elements(stressReport.2)] stressReport_likelihood;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        stressReport_gen[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] = kalman2_vector_rng(
            (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1))),
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            a11[plate_i__pl_1],
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls,
            q1,
            q2,
            r1,
            r2,
            s0[plate_i__pl_1],
            m0[plate_i__pl_1],
            10.0
        );
        stressReport_likelihood[plate_i__pl_1] = kalman2_lpdf(stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] | 
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            a11[plate_i__pl_1],
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls,
            q1,
            q2,
            r1,
            r2,
            s0[plate_i__pl_1],
            m0[plate_i__pl_1],
            10.0
        );
    }
}
julia
Turing unsupported for this BRM example

Turing backend: direct execution requires at least one observed likelihood

No latent state is a parameter any more: 38 dimensions on the 5-subject fixture, and that number does not grow with the number of occasions. The formula surface is unchanged in kind — per-subject parameters are still ordinary linear predictors with covariates and random effects; only the cell's likelihood was swapped.

3. The nonlinear, non-Gaussian model: an extended Kalman filter in the cell ​

The full model of section 1 is neither linear (softplus recovery) nor Gaussian (the binary indicator), so the marginalization is approximate: between observations the state mean and covariance are propagated by Euler substeps with the drift's Jacobian (an extended Kalman filter); at an observation the two continuous reports update the state exactly, and the binary indicator is integrated by Gauss–Hermite quadrature over the latent predictor, with a moment-matched state update.

julia
StanBlocks.@deffun begin
    l2pi()::real = 1.8378770664093453          # log(2 pi)
    ema_ekf_lpdfs(ys::vector[T], ym::vector[T], smoked::int[T], workload::vector[T], dt::vector[T],
            b0::real, bm::real, a12::real, a21::real, a22::real, cm::real, wls::real,
            q0::real, qw::real, diffm::real, diff21::real, l31::real, thr::real,
            mm_s::real, mm_m::real, r1::real, r2::real,
            ms0::real, mm0::real, t0sd1::real, t0sd2::real, t0z::real, nsub::int)::vector[T] = begin
        out::vector[T]; ms=ms0; mm=mm0; p11=t0sd1*t0sd1; p12=tanh(t0z)*t0sd1*t0sd2; p22=t0sd2*t0sd2
        for t in 1:T
            if t>1
                wl=workload[t]; h=dt[t]/nsub              # the diffusion reads the CURRENT row's input
                for st in 1:nsub
                    sp=log1p_exp(b0+bm*mm); sig=inv_logit(b0+bm*mm)
                    nms=ms+(-sp*ms+a12*mm)*h; nmm=mm+(a21*ms+a22*mm+cm)*h
                    f11=1-h*sp; f12=h*(a12-bm*sig*ms); f21=h*a21; f22=1+h*a22
                    sds=exp(q0+qw*wl); corr=tanh(diff21); qs=sds*sds*h; qc=corr*sds*diffm*h; qm=diffm*diffm*h
                    fp11=f11*p11+f12*p12; fp12=f11*p12+f12*p22; fp21=f21*p11+f22*p12; fp22=f21*p12+f22*p22
                    np11=fp11*f11+fp12*f12+qs; np12=fp11*f21+fp12*f22+qc; np22=fp21*f21+fp22*f22+qm
                    ms=nms; mm=nmm; p11=np11; p12=np12; p22=np22
                end
            end
            ms=ms+wls*workload[t]                    # workload acts as an impulse at the observation
            v1=ys[t]-(ms+mm_s); v2=ym[t]-(mm+mm_m); s11=p11+r1*r1; s12=p12; s22=p22+r2*r2
            det=s11*s22-s12*s12; si11=s22/det; si12=-s12/det; si22=s11/det
            quad=v1*(si11*v1+si12*v2)+v2*(si12*v1+si22*v2); lg=-0.5*(2*l2pi()+log(det)+quad)
            k11=p11*si11+p12*si12; k12=p11*si12+p12*si22; k21=p12*si11+p22*si12; k22=p12*si12+p22*si22
            ms=ms+k11*v1+k12*v2; mm=mm+k21*v1+k22*v2
            g11=(1-k11)*p11-k12*p12; g12=(1-k11)*p12-k12*p22; g21=-k21*p11+(1-k22)*p12; g22=-k21*p12+(1-k22)*p22
            p11=g11; p12=0.5*(g12+g21); p22=g22
            # Binary indicator: Gauss-Hermite (5 nodes) over eta ~ N(l31*stress+thr, l31^2*p11) with a
            # moment-matched state update -- as ctsem's `_binary_moments` does -- not linearised.
            etabar=l31*ms+thr; s2=l31*l31*p11+1e-12; sde=sqrt(s2); z0=0.0; z1=0.0; z2=0.0
            for i in 1:5
                xi=0.0; wi=0.5333333333333333
                if i==1; xi=-2.8569700138728056; wi=0.011257411327720689; end
                if i==2; xi=-1.3556261799742659; wi=0.22207592200561263; end
                if i==4; xi=1.3556261799742659; wi=0.22207592200561263; end
                if i==5; xi=2.8569700138728056; wi=0.011257411327720689; end
                eta=etabar+sde*xi; pr=1-inv_logit(eta)
                if smoked[t]==1; pr=inv_logit(eta); end
                z0=z0+wi*pr; z1=z1+wi*pr*eta; z2=z2+wi*pr*eta*eta
            end
            eeta=z1/z0; veta=z2/z0-eeta*eeta; kb1=p11*l31/s2; kb2=p12*l31/s2; shr=s2-veta
            ms=ms+kb1*(eeta-etabar); mm=mm+kb2*(eeta-etabar)
            n11=p11-kb1*kb1*shr; n12=p12-kb1*kb2*shr; n22=p22-kb2*kb2*shr
            p11=n11; p12=n12; p22=n22; out[t]=lg+log(z0)
        end
        out
    end
    @lhs @lpxf ema_ekf_lpdf(ys::vector[T], ym::vector[T], smoked::int[T],
            workload::vector[T], dt::vector[T],
            b0::real, bm::real, a12::real, a21::real, a22::real, cm::real, wls::real,
            q0::real, qw::real, diffm::real, diff21::real, l31::real, thr::real,
            mm_s::real, mm_m::real, r1::real, r2::real,
            ms0::real, mm0::real, t0sd1::real, t0sd2::real, t0z::real, nsub::int)::real = begin
        sum(ema_ekf_lpdfs(ys, ym, smoked, workload, dt, b0, bm, a12, a21, a22, cm, wls, q0, qw, diffm, diff21,
                          l31, thr, mm_s, mm_m, r1, r2, ms0, mm0, t0sd1, t0sd2, t0z, nsub))
    end
    ema_ekf_rng(vector[T], ym::vector[T], smoked::int[T], workload::vector[T], dt::vector[T],
            b0::real, bm::real, a12::real, a21::real, a22::real, cm::real, wls::real,
            q0::real, qw::real, diffm::real, diff21::real, l31::real, thr::real,
            mm_s::real, mm_m::real, r1::real, r2::real,
            ms0::real, mm0::real, t0sd1::real, t0sd2::real, t0z::real, nsub::int)::vector[T] = begin
        out::vector[T]; z01=normal_rng(0.,1.); z02=normal_rng(0.,1.); r0=tanh(t0z)
        s=ms0+t0sd1*z01; m=mm0+t0sd2*(r0*z01+sqrt(1-r0*r0)*z02)
        for t in 1:T
            if t>1
                wl=workload[t]; h=dt[t]/nsub; corr=tanh(diff21)
                for st in 1:nsub
                    sds=exp(q0+qw*wl); z1=normal_rng(0.,1.); z2=corr*z1+sqrt(1-corr*corr)*normal_rng(0.,1.)
                    ds=(-log1p_exp(b0+bm*m)*s+a12*m)*h+sds*sqrt(h)*z1
                    dm=(a21*s+a22*m+cm)*h+diffm*sqrt(h)*z2
                    s=s+ds; m=m+dm
                end
            end
            s=s+wls*workload[t]                       # impulse at the observation
            out[t]=normal_rng(s+mm_s,r1)
        end
        out
    end
end
brm-comparison
Hierarchical EMA model, states EKF-marginalized per subject
julia
function ema_kernel_ekf_model(data = ema_ekf_fixture())
    @brm data begin
        # shared drift / coupling / observation parameters
        bm    ~ Normal(0, 0.5)                  # mood -> stress recovery modulation
        a12   ~ Normal(0, 0.5)                  # mood -> stress
        a21   ~ Normal(0, 0.5)                  # stress -> mood
        a22   ~ Normal(-0.5, 0.3)               # mood self-decay
        qw    ~ Normal(0, 0.5)                  # workload -> stress process noise
        diffm ~ Exponential(1.0)                # mood process-noise sd
        diff21 ~ Normal(0, 0.5)                 # process-noise correlation (fisher-z)
        l31   ~ Normal(0, 1)                    # stress -> smoking loading
        thr   ~ Normal(0, 1)                    # smoking threshold
        mm_s  ~ Normal(0, 0.5)                  # manifest mean, stressReport
        mm_m  ~ Normal(0, 0.5)                  # manifest mean, moodReport
        r1    ~ Exponential(1.0)                # measurement sd, stressReport
        r2    ~ Exponential(1.0)                # measurement sd, moodReport
        s0    ~ Normal(0, 1)                    # initial stress mean
        m0    ~ Normal(0, 1)                    # initial mood mean
        t0sd1 ~ Exponential(1.0)                # initial covariance: stress sd
        t0sd2 ~ Exponential(1.0)                #                     mood sd
        t0z   ~ Normal(0, 0.5)                  #                     fisher-z correlation

        # the four subject-varying parameters share ONE correlated random-effect block
        b0  ~ 1 + age + treatment + (1 | p | subject)   # stress-recovery baseline
        q0  ~ 1 +                   (1 | p | subject)   # stress process-noise baseline
        cm  ~ 1 +       treatment + (1 | p | subject)   # mood intercept (cint_mood)
        wls ~ 1 +                   (1 | p | subject)   # workload -> stress impulse (wl_stress)

        # per subject: integrate out THIS subject's entire latent path with the EKF
        pred ~ kernel(dt, workload, stressReport, moodReport, smoked,
                      b0, q0, cm, wls) do dti, wli, ys, ym, smk, lb0, lq0, lcm, lwls
            ys ~ ema_ekf(ym, smk, wli, dti,
                         lb0, bm, a12, a21, a22, lcm, lwls, lq0, qw, diffm, diff21, l31, thr,
                         mm_s, mm_m, r1, r2, s0, m0, t0sd1, t0sd2, t0z, 4)
            ys
        end
    end
end
julia
BRMI:
  bm ~ Normal(0, 0.5)
  a12 ~ Normal(0, 0.5)
  a21 ~ Normal(0, 0.5)
  a22 ~ Normal(-0.5, 0.3)
  qw ~ Normal(0, 0.5)
  diffm ~ Exponential(1.0)
  diff21 ~ Normal(0, 0.5)
  l31 ~ Normal(0, 1)
  thr ~ Normal(0, 1)
  mm_s ~ Normal(0, 0.5)
  mm_m ~ Normal(0, 0.5)
  r1 ~ Exponential(1.0)
  r2 ~ Exponential(1.0)
  s0 ~ Normal(0, 1)
  m0 ~ Normal(0, 1)
  t0sd1 ~ Exponential(1.0)
  t0sd2 ~ Exponential(1.0)
  t0z ~ Normal(0, 0.5)
  age: data (eltype=Float64, n=6)
  treatment: data (eltype=Float64, n=6)
  subject: data (eltype=String, n=6)
  b0 ~ 1 + age + treatment + (1 | p | subject)
  q0 ~ 1 + (1 | p | subject)
  cm ~ 1 + treatment + (1 | p | subject)
  wls ~ 1 + (1 | p | subject)
  dt: data (eltype=Vector{Float64}, n=6)
  workload: data (eltype=Vector{Float64}, n=6)
  stressReport: data (eltype=Vector{Float64}, n=6)
  moodReport: data (eltype=Vector{Float64}, n=6)
  smoked: data (eltype=Vector{Int64}, n=6)
  pred ~ kernel((dti, wli, ys, ym, smk, lb0, lq0, lcm, lwls)->begin
        #= brm-docs-example.jl:32 =#
        ys ~ ema_ekf(ym, smk, wli, dti, lb0, bm, a12, a21, a22, lcm, lwls, lq0, qw, diffm, diff21, l31, thr, mm_s, mm_m, r1, r2, s0, m0, t0sd1, t0sd2, t0z, 4)
        #= brm-docs-example.jl:35 =#
        ys
    end, dt, workload, stressReport, moodReport, smoked, b0, q0, cm, wls)
julia
SBBRMI with data keys = [:age, :dt, :kernel_nsub_pred, :moodReport, :n_subject, :n_terms_p_subject, :smoked, :stressReport, :subject_idx, :treatment, :workload]
emitted @slic body:
begin
    b_p_subject ~ ranef_correlated_draws(; group_idx = subject_idx, n_groups = n_subject, n_terms = n_terms_p_subject)
    bm ~ normal(0, 0.5)
    a12 ~ normal(0, 0.5)
    a21 ~ normal(0, 0.5)
    a22 ~ normal(-0.5, 0.3)
    qw ~ normal(0, 0.5)
    diffm ~ exponential(1.0 ./ 1.0)
    diff21 ~ normal(0, 0.5)
    l31 ~ normal(0, 1)
    thr ~ normal(0, 1)
    mm_s ~ normal(0, 0.5)
    mm_m ~ normal(0, 0.5)
    r1 ~ exponential(1.0 ./ 1.0)
    r2 ~ exponential(1.0 ./ 1.0)
    s0 ~ normal(0, 1)
    m0 ~ normal(0, 1)
    t0sd1 ~ exponential(1.0 ./ 1.0)
    t0sd2 ~ exponential(1.0 ./ 1.0)
    t0z ~ normal(0, 0.5)
    X_b0 = hcat(rep_vector(1.0, num_elements(age)), age, treatment)
    pop_b0 ~ popefs(; X = X_b0)
    r_b0_p_subject = b_p_subject[subject_idx, 1]
    b0 = pop_b0 + r_b0_p_subject
    X_q0 = hcat(rep_vector(1.0, num_elements(subject_idx)))
    pop_q0 ~ popefs(; X = X_q0)
    r_q0_p_subject = b_p_subject[subject_idx, 2]
    q0 = pop_q0 + r_q0_p_subject
    X_cm = hcat(rep_vector(1.0, num_elements(treatment)), treatment)
    pop_cm ~ popefs(; X = X_cm)
    r_cm_p_subject = b_p_subject[subject_idx, 3]
    cm = pop_cm + r_cm_p_subject
    X_wls = hcat(rep_vector(1.0, num_elements(subject_idx)))
    pop_wls ~ popefs(; X = X_wls)
    r_wls_p_subject = b_p_subject[subject_idx, 4]
    wls = pop_wls + r_wls_p_subject
    pred ~ plate(dt, workload, stressReport, moodReport, smoked, b0, q0, cm, wls; outer = (kernel_nsub_pred,)) do dti, wli, ys, ym, smk, lb0, lq0, lcm, lwls
            #= brm-docs-example.jl:32 =#
            ys ~ ema_ekf(ym, smk, wli, dti, lb0, bm, a12, a21, a22, lcm, lwls, lq0, qw, diffm, diff21, l31, thr, mm_s, mm_m, r1, r2, s0, m0, t0sd1, t0sd2, t0z, 4)
            #= brm-docs-example.jl:35 =#
            ys
        end
end
stan
functions {
matrix hcat(vector x, vector y, vector z) {
    return hcat(hcat(x, y), z);
}
matrix hcat(
    matrix x,
    vector y
) {
    int m = dims(x)[1];
    int n = dims(x)[2];
    if (dims(y)[1] != m) reject("hcat: dim mismatch — `y` dim 1 (= ", dims(y)[1], ") does not match `m` (= ", m, "), inferred from `x` dim 1. `m` sizes: `x` dim 1 (= ", dims(x)[1], "), `y` dim 1 (= ", dims(y)[1], ").");
    return append_col(x, y);
}
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);
}
matrix hcat(vector x) {
    int n = dims(x)[1];
    return to_matrix(x, n, 1);
}
int ragged_end(array[] int ends, int i) {
    return ends[i];
}
int ragged_start(
    array[] int ends,
    int i
) {
    if((i == 1)) {
        return 1;
    } else {
        return (1 + ends[(i - 1)]);
    }
}
real ema_ekf_lpdf(
    vector ys,
    vector ym,
    array[] int smoked,
    vector workload,
    vector dt,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q0,
    real qw,
    real diffm,
    real diff21,
    real l31,
    real thr,
    real mm_s,
    real mm_m,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real t0sd1,
    real t0sd2,
    real t0z,
    int nsub
) {
    int T = dims(ys)[1];
    if (dims(ym)[1] != T) reject("ema_ekf_lpdf: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(smoked)[1] != T) reject("ema_ekf_lpdf: dim mismatch — `smoked` dim 1 (= ", dims(smoked)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(workload)[1] != T) reject("ema_ekf_lpdf: dim mismatch — `workload` dim 1 (= ", dims(workload)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("ema_ekf_lpdf: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    return sum(
        ema_ekf_lpdfs(
            ys,
            ym,
            smoked,
            workload,
            dt,
            b0,
            bm,
            a12,
            a21,
            a22,
            cm,
            wls,
            q0,
            qw,
            diffm,
            diff21,
            l31,
            thr,
            mm_s,
            mm_m,
            r1,
            r2,
            ms0,
            mm0,
            t0sd1,
            t0sd2,
            t0z,
            nsub
        )
    );
}
vector ema_ekf_lpdfs(
    vector ys,
    vector ym,
    array[] int smoked,
    vector workload,
    vector dt,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q0,
    real qw,
    real diffm,
    real diff21,
    real l31,
    real thr,
    real mm_s,
    real mm_m,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real t0sd1,
    real t0sd2,
    real t0z,
    int nsub
) {
    int T = dims(ys)[1];
    if (dims(ym)[1] != T) reject("ema_ekf_lpdfs: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(smoked)[1] != T) reject("ema_ekf_lpdfs: dim mismatch — `smoked` dim 1 (= ", dims(smoked)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(workload)[1] != T) reject("ema_ekf_lpdfs: dim mismatch — `workload` dim 1 (= ", dims(workload)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("ema_ekf_lpdfs: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    vector[T] out;
    real ms = ms0;
    real mm = mm0;
    real p11 = (t0sd1 * t0sd1);
    real p12 = (tanh(t0z) * t0sd1 * t0sd2);
    real p22 = (t0sd2 * t0sd2);
    for(t in 1:T) {
        if((t > 1)) {
            real wl = workload[t];
            real h = (dt[t] / nsub);
            for(st in 1:nsub) {
                real sp = log1p_exp((b0 + (bm * mm)));
                real sig = inv_logit((b0 + (bm * mm)));
                real nms = (ms + ((((-sp) * ms) + (a12 * mm)) * h));
                real nmm = (mm + (((a21 * ms) + (a22 * mm) + cm) * h));
                real f11 = (1 - (h * sp));
                real f12 = (h * (a12 - (bm * sig * ms)));
                real f21 = (h * a21);
                real f22 = (1 + (h * a22));
                real sds = exp((q0 + (qw * wl)));
                real corr = tanh(diff21);
                real qs = (sds * sds * h);
                real qc = (corr * sds * diffm * h);
                real qm = (diffm * diffm * h);
                real fp11 = ((f11 * p11) + (f12 * p12));
                real fp12 = ((f11 * p12) + (f12 * p22));
                real fp21 = ((f21 * p11) + (f22 * p12));
                real fp22 = ((f21 * p12) + (f22 * p22));
                real np11 = ((fp11 * f11) + (fp12 * f12) + qs);
                real np12 = ((fp11 * f21) + (fp12 * f22) + qc);
                real np22 = ((fp21 * f21) + (fp22 * f22) + qm);
                ms = nms;
                mm = nmm;
                p11 = np11;
                p12 = np12;
                p22 = np22;
            }
        }
        ms = (ms + (wls * workload[t]));
        real v1 = (ys[t] - (ms + mm_s));
        real v2 = (ym[t] - (mm + mm_m));
        real s11 = (p11 + (r1 * r1));
        real s12 = p12;
        real s22 = (p22 + (r2 * r2));
        real det = ((s11 * s22) - (s12 * s12));
        real si11 = (s22 / det);
        real si12 = ((-s12) / det);
        real si22 = (s11 / det);
        real quad = ((v1 * ((si11 * v1) + (si12 * v2))) + (v2 * ((si12 * v1) + (si22 * v2))));
        real lg = (-0.5 * ((2 * l2pi()) + log(det) + quad));
        real k11 = ((p11 * si11) + (p12 * si12));
        real k12 = ((p11 * si12) + (p12 * si22));
        real k21 = ((p12 * si11) + (p22 * si12));
        real k22 = ((p12 * si12) + (p22 * si22));
        ms = (ms + (k11 * v1) + (k12 * v2));
        mm = (mm + (k21 * v1) + (k22 * v2));
        real g11 = (((1 - k11) * p11) - (k12 * p12));
        real g12 = (((1 - k11) * p12) - (k12 * p22));
        real g21 = (((-k21) * p11) + ((1 - k22) * p12));
        real g22 = (((-k21) * p12) + ((1 - k22) * p22));
        p11 = g11;
        p12 = (0.5 * (g12 + g21));
        p22 = g22;
        real etabar = ((l31 * ms) + thr);
        real s2 = ((l31 * l31 * p11) + 1.0e-12);
        real sde = sqrt(s2);
        real z0 = 0.0;
        real z1 = 0.0;
        real z2 = 0.0;
        for(i in 1:5) {
            real xi = 0.0;
            real wi = 0.5333333333333333;
            if((i == 1)) {
                xi = -2.8569700138728056;
                wi = 0.01125741132772069;
            }
            if((i == 2)) {
                xi = -1.355626179974266;
                wi = 0.22207592200561263;
            }
            if((i == 4)) {
                xi = 1.355626179974266;
                wi = 0.22207592200561263;
            }
            if((i == 5)) {
                xi = 2.8569700138728056;
                wi = 0.01125741132772069;
            }
            real eta = (etabar + (sde * xi));
            real pr = (1 - inv_logit(eta));
            if((smoked[t] == 1)) {
                pr = inv_logit(eta);
            }
            z0 = (z0 + (wi * pr));
            z1 = (z1 + (wi * pr * eta));
            z2 = (z2 + (wi * pr * eta * eta));
        }
        real eeta = (z1 / z0);
        real veta = ((z2 / z0) - (eeta * eeta));
        real kb1 = ((p11 * l31) / s2);
        real kb2 = ((p12 * l31) / s2);
        real shr = (s2 - veta);
        ms = (ms + (kb1 * (eeta - etabar)));
        mm = (mm + (kb2 * (eeta - etabar)));
        real n11 = (p11 - (kb1 * kb1 * shr));
        real n12 = (p12 - (kb1 * kb2 * shr));
        real n22 = (p22 - (kb2 * kb2 * shr));
        p11 = n11;
        p12 = n12;
        p22 = n22;
        out[t] = (lg + log(z0));
    }
    return out;
}
real l2pi() {
    return 1.8378770664093453;
}
vector ema_ekf_vector_rng(
    int anontok__1,
    vector ym,
    array[] int smoked,
    vector workload,
    vector dt,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cm,
    real wls,
    real q0,
    real qw,
    real diffm,
    real diff21,
    real l31,
    real thr,
    real mm_s,
    real mm_m,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real t0sd1,
    real t0sd2,
    real t0z,
    int nsub
) {
    int T = anontok__1;
    if (dims(ym)[1] != T) reject("ema_ekf_rng: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(smoked)[1] != T) reject("ema_ekf_rng: dim mismatch — `smoked` dim 1 (= ", dims(smoked)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(workload)[1] != T) reject("ema_ekf_rng: dim mismatch — `workload` dim 1 (= ", dims(workload)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("ema_ekf_rng: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `workload` dim 1 (= ", dims(workload)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    vector[T] out;
    real z01 = normal_rng(0.0, 1.0);
    real z02 = normal_rng(0.0, 1.0);
    real r0 = tanh(t0z);
    real s = (ms0 + (t0sd1 * z01));
    real m = (mm0 + (t0sd2 * ((r0 * z01) + (sqrt((1 - (r0 * r0))) * z02))));
    for(t in 1:T) {
        if((t > 1)) {
            real wl = workload[t];
            real h = (dt[t] / nsub);
            real corr = tanh(diff21);
            for(st in 1:nsub) {
                real sds = exp((q0 + (qw * wl)));
                real z1 = normal_rng(0.0, 1.0);
                real z2 = ((corr * z1) + (sqrt((1 - (corr * corr))) * normal_rng(0.0, 1.0)));
                real ds = (((((-log1p_exp((b0 + (bm * m)))) * s) + (a12 * m)) * h) + (sds * sqrt(h) * z1));
                real dm = ((((a21 * s) + (a22 * m) + cm) * h) + (diffm * sqrt(h) * z2));
                s = (s + ds);
                m = (m + dm);
            }
        }
        s = (s + (wls * workload[t]));
        out[t] = normal_rng((s + mm_s), r1);
    }
    return out;
}
}
data {
    int n_terms_p_subject;
    int n_subject;
    int treatment_n;
    int age_n;
    vector[age_n] age;
    vector[treatment_n] treatment;
    int subject_idx_n;
    array[subject_idx_n] int subject_idx;
    int kernel_nsub_pred;
    int stressReport_ends_n;
    int stressReport_mem_n;
    tuple(vector[stressReport_mem_n], array[stressReport_ends_n] int) stressReport;
    int dt_ends_n;
    int dt_mem_n;
    tuple(vector[dt_mem_n], array[dt_ends_n] int) dt;
    int moodReport_ends_n;
    int moodReport_mem_n;
    tuple(vector[moodReport_mem_n], array[moodReport_ends_n] int) moodReport;
    int smoked_ends_n;
    int smoked_mem_n;
    tuple(array[smoked_mem_n] int, array[smoked_ends_n] int) smoked;
    int workload_ends_n;
    int workload_mem_n;
    tuple(vector[workload_mem_n], array[workload_ends_n] int) workload;
}
transformed data {
    matrix[treatment_n, (2 + 1)] X_b0 = hcat(rep_vector(1.0, num_elements(age)), age, treatment);
    int pop_b0_n_covariates = (2 + 1);
    matrix[num_elements(subject_idx), 1] X_q0 = hcat(rep_vector(1.0, num_elements(subject_idx)));
    int pop_q0_n_covariates = 1;
    matrix[treatment_n, 2] X_cm = hcat(rep_vector(1.0, num_elements(treatment)), treatment);
    int pop_cm_n_covariates = 2;
    matrix[num_elements(subject_idx), 1] X_wls = hcat(rep_vector(1.0, num_elements(subject_idx)));
    int pop_wls_n_covariates = 1;
    array[kernel_nsub_pred] int pred__pl_len_1;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        pred__pl_len_1[plate_i__pl_1] = (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1)));
    }
    array[kernel_nsub_pred] int pred__pl_end_1 = cumulative_sum(pred__pl_len_1);
    vector[sum(pred__pl_len_1)] pred__pl_mem_1;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        pred__pl_mem_1[
            ragged_start(pred__pl_end_1, plate_i__pl_1):ragged_end(pred__pl_end_1, plate_i__pl_1)
        ] = stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ];
    }
}
parameters {
    cholesky_factor_corr[n_terms_p_subject] b_p_subject_L;
    vector<lower=0.0>[n_terms_p_subject] b_p_subject_tau;
    vector[(n_terms_p_subject * n_subject)] b_p_subject_z_flat;
    real bm;
    real a12;
    real a21;
    real a22;
    real qw;
    real<lower=0.0> diffm;
    real diff21;
    real l31;
    real thr;
    real mm_s;
    real mm_m;
    real<lower=0.0> r1;
    real<lower=0.0> r2;
    real s0;
    real m0;
    real<lower=0.0> t0sd1;
    real<lower=0.0> t0sd2;
    real t0z;
    vector[pop_b0_n_covariates] pop_b0_beta_pop;
    vector[pop_q0_n_covariates] pop_q0_beta_pop;
    vector[pop_cm_n_covariates] pop_cm_beta_pop;
    vector[pop_wls_n_covariates] pop_wls_beta_pop;
}
transformed parameters {
    matrix[n_terms_p_subject, n_subject] b_p_subject_z = to_matrix(b_p_subject_z_flat, n_terms_p_subject, n_subject);
    matrix[n_subject, n_terms_p_subject] b_p_subject = ((diag_pre_multiply(b_p_subject_tau, b_p_subject_L) * b_p_subject_z)');
    vector[treatment_n] pop_b0 = (X_b0 * pop_b0_beta_pop);
    vector[subject_idx_n] r_b0_p_subject = b_p_subject[subject_idx, 1];
    vector[treatment_n] b0 = (pop_b0 + r_b0_p_subject);
    vector[num_elements(subject_idx)] pop_q0 = (X_q0 * pop_q0_beta_pop);
    vector[subject_idx_n] r_q0_p_subject = b_p_subject[subject_idx, 2];
    vector[num_elements(subject_idx)] q0 = (pop_q0 + r_q0_p_subject);
    vector[treatment_n] pop_cm = (X_cm * pop_cm_beta_pop);
    vector[subject_idx_n] r_cm_p_subject = b_p_subject[subject_idx, 3];
    vector[treatment_n] cm = (pop_cm + r_cm_p_subject);
    vector[num_elements(subject_idx)] pop_wls = (X_wls * pop_wls_beta_pop);
    vector[subject_idx_n] r_wls_p_subject = b_p_subject[subject_idx, 4];
    vector[num_elements(subject_idx)] wls = (pop_wls + r_wls_p_subject);
}
model {
    b_p_subject_L ~ lkj_corr_cholesky(1.0);
    b_p_subject_tau ~ std_normal();
    b_p_subject_z_flat ~ std_normal();
    bm ~ normal(0, 0.5);
    a12 ~ normal(0, 0.5);
    a21 ~ normal(0, 0.5);
    a22 ~ normal(-0.5, 0.3);
    qw ~ normal(0, 0.5);
    diffm ~ exponential((1.0 ./ 1.0));
    diff21 ~ normal(0, 0.5);
    l31 ~ normal(0, 1);
    thr ~ normal(0, 1);
    mm_s ~ normal(0, 0.5);
    mm_m ~ normal(0, 0.5);
    r1 ~ exponential((1.0 ./ 1.0));
    r2 ~ exponential((1.0 ./ 1.0));
    s0 ~ normal(0, 1);
    m0 ~ normal(0, 1);
    t0sd1 ~ exponential((1.0 ./ 1.0));
    t0sd2 ~ exponential((1.0 ./ 1.0));
    t0z ~ normal(0, 0.5);
    pop_b0_beta_pop ~ std_normal();
    pop_q0_beta_pop ~ std_normal();
    pop_cm_beta_pop ~ std_normal();
    pop_wls_beta_pop ~ std_normal();
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] ~ ema_ekf(
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            b0[plate_i__pl_1],
            bm,
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls[plate_i__pl_1],
            q0[plate_i__pl_1],
            qw,
            diffm,
            diff21,
            l31,
            thr,
            mm_s,
            mm_m,
            r1,
            r2,
            s0,
            m0,
            t0sd1,
            t0sd2,
            t0z,
            4
        );
    }
}
generated quantities {
    vector[num_elements(stressReport.1)] stressReport_gen;
    vector[num_elements(stressReport.2)] stressReport_likelihood;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        stressReport_gen[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] = ema_ekf_vector_rng(
            (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1))),
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            b0[plate_i__pl_1],
            bm,
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls[plate_i__pl_1],
            q0[plate_i__pl_1],
            qw,
            diffm,
            diff21,
            l31,
            thr,
            mm_s,
            mm_m,
            r1,
            r2,
            s0,
            m0,
            t0sd1,
            t0sd2,
            t0z,
            4
        );
        stressReport_likelihood[plate_i__pl_1] = ema_ekf_lpdf(stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] | 
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)],
            workload.1[ragged_start(workload.2, plate_i__pl_1):ragged_end(workload.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            b0[plate_i__pl_1],
            bm,
            a12,
            a21,
            a22,
            cm[plate_i__pl_1],
            wls[plate_i__pl_1],
            q0[plate_i__pl_1],
            qw,
            diffm,
            diff21,
            l31,
            thr,
            mm_s,
            mm_m,
            r1,
            r2,
            s0,
            m0,
            t0sd1,
            t0sd2,
            t0z,
            4
        );
    }
}
julia
Turing unsupported for this BRM example

Turing backend: direct execution requires at least one observed likelihood

This is the same population model as in section 1 — the same covariates, the same correlated (1 | p | subject) block — at 59 dimensions instead of 190, and it stays at 59 however long the series get.

4. Dynamics that depend on the latent state ​

In the second demonstration model three cells of the system matrices are functions of the latent state: stress recovers faster in good mood, stress is more volatile in good mood, and the two shocks couple more tightly the more stressed the person is,

A11=−softplus(b0+bmmood),G11=exp⁡(q0+q1mood),corr(dW1,dW2)=tanh⁡(czstress).

So the process-noise covariance depends on the very states being integrated out. The filter's predict step comes in two orders: first-order (drift Jacobian and noise covariance evaluated at the filtered mean), or moment-matched — the Euler map and the noise covariance averaged over the current state uncertainty with a Gauss–Hermite rule.

julia
StanBlocks.@deffun begin
    l2pi()::real = 1.8378770664093453          # log(2 pi)
    # probabilists' Gauss-Hermite node / weight i of a K-point rule (K = 3 or 5)
    ghx(i::int, K::int)::real = begin
        x=0.0
        if K==3; x=(i-2)*1.7320508075688772; end
        if K==5
            if i==1; x=-2.8569700138728056; end
            if i==2; x=-1.3556261799742659; end
            if i==4; x=1.3556261799742659; end
            if i==5; x=2.8569700138728056; end
        end
        x
    end
    ghw(i::int, K::int)::real = begin
        w=0.6666666666666666
        if K==3
            if i!=2; w=0.16666666666666666; end
        end
        if K==5
            w=0.5333333333333333
            if i==1 || i==5; w=0.011257411327720689; end
            if i==2 || i==4; w=0.22207592200561263; end
        end
        w
    end
    ema_sd_lpdfs(ys::vector[T], ym::vector[T], smoked::int[T], dt::vector[T],
            b0::real, bm::real, a12::real, a21::real, a22::real, cintm::real,
            qd0::real, qd1::real, cz::real, sdm::real,
            l31::real, thr::real, r1::real, r2::real,
            ms0::real, mm0::real, t0sd1::real, t0sd2::real, t0z::real, nsub::int, gh::int)::vector[T] = begin
        out::vector[T]; ms=ms0; mm=mm0; p11=t0sd1*t0sd1; p12=tanh(t0z)*t0sd1*t0sd2; p22=t0sd2*t0sd2
        for t in 1:T
            if t>1
                h=dt[t]/nsub                       # SUBSTEPPED continuous-time predict
                for st in 1:nsub
                    if gh>0
                        l11=sqrt(p11); l21=p12/l11; l22=sqrt(p22-l21*l21+1e-12)
                        ey1=0.0; ey2=0.0; c11=0.0; c12=0.0; c22=0.0; q11=0.0; q12=0.0
                        for a in 1:gh
                            xa=ghx(a,gh)
                            for b in 1:gh
                                xb=ghx(b,gh); w=ghw(a,gh)*ghw(b,gh)
                                xs=ms+l11*xa; xm=mm+l21*xa+l22*xb
                                y1=xs+h*(-log1p_exp(b0+bm*xm)*xs+a12*xm)
                                y2=xm+h*(a21*xs+a22*xm+cintm)
                                sdx=exp(qd0+qd1*xm)            # stress sd depends on MOOD
                                ey1=ey1+w*y1; ey2=ey2+w*y2
                                c11=c11+w*y1*y1; c12=c12+w*y1*y2; c22=c22+w*y2*y2
                                q11=q11+w*sdx*sdx
                                q12=q12+w*tanh(cz*xs)*sdx*sdm  # shock corr depends on STRESS
                            end
                        end
                        p11=c11-ey1*ey1+h*q11; p12=c12-ey1*ey2+h*q12; p22=c22-ey2*ey2+h*sdm*sdm
                        ms=ey1; mm=ey2
                    else
                        sp=log1p_exp(b0+bm*mm); sig=inv_logit(b0+bm*mm)
                        nms=ms+(-sp*ms+a12*mm)*h; nmm=mm+(a21*ms+a22*mm+cintm)*h
                        f11=1-h*sp; f12=h*(a12-bm*sig*ms); f21=h*a21; f22=1+h*a22
                        sds=exp(qd0+qd1*mm); corr=tanh(cz*ms)
                        qs=sds*sds*h; qc=corr*sds*sdm*h; qm=sdm*sdm*h
                        fp11=f11*p11+f12*p12; fp12=f11*p12+f12*p22; fp21=f21*p11+f22*p12; fp22=f21*p12+f22*p22
                        np11=fp11*f11+fp12*f12+qs; np12=fp11*f21+fp12*f22+qc; np22=fp21*f21+fp22*f22+qm
                        ms=nms; mm=nmm; p11=np11; p12=np12; p22=np22
                    end
                end
            end
            # Gaussian update (2 continuous indicators, loadings [1;1]); r1, r2 are measurement sds
            v1=ys[t]-ms; v2=ym[t]-mm; s11=p11+r1*r1; s12=p12; s22=p22+r2*r2
            det=s11*s22-s12*s12; si11=s22/det; si12=-s12/det; si22=s11/det
            quad=v1*(si11*v1+si12*v2)+v2*(si12*v1+si22*v2); lg=-0.5*(2*l2pi()+log(det)+quad)
            k11=p11*si11+p12*si12; k12=p11*si12+p12*si22; k21=p12*si11+p22*si12; k22=p12*si12+p22*si22
            ms=ms+k11*v1+k12*v2; mm=mm+k21*v1+k22*v2
            g11=(1-k11)*p11-k12*p12; g12=(1-k11)*p12-k12*p22; g21=-k21*p11+(1-k22)*p12; g22=-k21*p12+(1-k22)*p22
            p11=g11; p12=0.5*(g12+g21); p22=g22
            # Binary indicator: Gauss-Hermite (5 nodes) over eta ~ N(l31*stress+thr, l31^2*p11)
            etabar=l31*ms+thr; s2=l31*l31*p11+1e-12; sde=sqrt(s2); z0=0.0; z1=0.0; z2=0.0
            for i in 1:5
                xi=ghx(i,5); wi=ghw(i,5)
                eta=etabar+sde*xi; pr=1-inv_logit(eta)
                if smoked[t]==1; pr=inv_logit(eta); end
                z0=z0+wi*pr; z1=z1+wi*pr*eta; z2=z2+wi*pr*eta*eta
            end
            eeta=z1/z0; veta=z2/z0-eeta*eeta; kb1=p11*l31/s2; kb2=p12*l31/s2; shr=s2-veta
            ms=ms+kb1*(eeta-etabar); mm=mm+kb2*(eeta-etabar)
            n11=p11-kb1*kb1*shr; n12=p12-kb1*kb2*shr; n22=p22-kb2*kb2*shr
            p11=n11; p12=n12; p22=n22
            out[t]=lg+log(z0)
        end
        out
    end
    @lhs @lpxf ema_sd_lpdf(ys::vector[T], ym::vector[T], smoked::int[T], dt::vector[T],
            b0::real, bm::real, a12::real, a21::real, a22::real, cintm::real,
            qd0::real, qd1::real, cz::real, sdm::real,
            l31::real, thr::real, r1::real, r2::real,
            ms0::real, mm0::real, t0sd1::real, t0sd2::real, t0z::real, nsub::int, gh::int)::real = begin
        sum(ema_sd_lpdfs(ys, ym, smoked, dt, b0, bm, a12, a21, a22, cintm, qd0, qd1, cz, sdm, l31, thr, r1, r2, ms0, mm0, t0sd1, t0sd2, t0z, nsub, gh))
    end
    ema_sd_rng(vector[T], ym::vector[T], smoked::int[T], dt::vector[T],
            b0::real, bm::real, a12::real, a21::real, a22::real, cintm::real,
            qd0::real, qd1::real, cz::real, sdm::real,
            l31::real, thr::real, r1::real, r2::real,
            ms0::real, mm0::real, t0sd1::real, t0sd2::real, t0z::real, nsub::int, gh::int)::vector[T] = begin
        out::vector[T]; z01=normal_rng(0.,1.); z02=normal_rng(0.,1.); r0=tanh(t0z)
        s=ms0+t0sd1*z01; m=mm0+t0sd2*(r0*z01+sqrt(1-r0*r0)*z02)
        for t in 1:T
            if t>1
                h=dt[t]/nsub
                for st in 1:nsub
                    sds=exp(qd0+qd1*m); corr=tanh(cz*s)
                    zs=normal_rng(0.,1.); zc=normal_rng(0.,1.); z2=corr*zs+sqrt(1-corr*corr)*zc
                    ds=(-log1p_exp(b0+bm*m)*s+a12*m)*h+sds*sqrt(h)*zs
                    dm=(a21*s+a22*m+cintm)*h+sdm*sqrt(h)*z2
                    s=s+ds; m=m+dm
                end
            end
            out[t]=normal_rng(s,r1)
        end
        out
    end
end

All parameters are shared across subjects here, so there is no random-effect grouping; the kernel takes the subject count from the pre-grouped (vector-of-vectors) columns.

brm-comparison
State-dependent drift and diffusion, states marginalized per subject
julia
function ema_state_dependent_model(data = with_filter(ema_state_dependent_fixture()))
    @brm data begin
        b0    ~ Normal(0.5, 0.5)               # stress recovery: softplus offset
        bm    ~ Normal(0.4, 0.5)               #                  modulation by mood
        a12   ~ Normal(-0.25, 0.5)             # mood -> stress
        a21   ~ Normal(-0.30, 0.5)             # stress -> mood
        a22   ~ Normal(-0.60, 0.3)             # mood self-decay
        cintm ~ Normal(0.3, 0.5)               # mood intercept
        qd0   ~ Normal(-0.2, 0.5)              # stress log-sd: offset
        qd1   ~ Normal(0.3, 0.5)               #                dependence on MOOD
        cz    ~ Normal(0.7, 0.5)               # shock correlation: dependence on STRESS
        sdm   ~ Exponential(1.0)               # mood diffusion sd
        l31   ~ Normal(1.2, 0.5)               # binary indicator: loading on stress
        thr   ~ Normal(-1.0, 0.5)              #                   threshold
        r1    ~ Exponential(1.0)               # measurement sd, stressReport
        r2    ~ Exponential(1.0)               # measurement sd, moodReport
        s0    ~ Normal(0.0, 1.0)               # initial stress mean
        m0    ~ Normal(0.5, 1.0)               # initial mood mean
        t0sd1 ~ Exponential(1.0)               # initial covariance: stress sd
        t0sd2 ~ Exponential(1.0)               #                     mood sd
        t0z   ~ Normal(0.0, 0.5)               #                     fisher-z correlation
        pred ~ kernel(dt, stressReport, moodReport, smoked) do dti, ys, ym, smk
            ys ~ ema_sd(ym, smk, dti, b0, bm, a12, a21, a22, cintm, qd0, qd1, cz, sdm,
                        l31, thr, r1, r2, s0, m0, t0sd1, t0sd2, t0z, nsub, gh)   # filter precision: DATA
            ys
        end
    end
end
julia
BRMI:
  b0 ~ Normal(0.5, 0.5)
  bm ~ Normal(0.4, 0.5)
  a12 ~ Normal(-0.25, 0.5)
  a21 ~ Normal(-0.3, 0.5)
  a22 ~ Normal(-0.6, 0.3)
  cintm ~ Normal(0.3, 0.5)
  qd0 ~ Normal(-0.2, 0.5)
  qd1 ~ Normal(0.3, 0.5)
  cz ~ Normal(0.7, 0.5)
  sdm ~ Exponential(1.0)
  l31 ~ Normal(1.2, 0.5)
  thr ~ Normal(-1.0, 0.5)
  r1 ~ Exponential(1.0)
  r2 ~ Exponential(1.0)
  s0 ~ Normal(0.0, 1.0)
  m0 ~ Normal(0.5, 1.0)
  t0sd1 ~ Exponential(1.0)
  t0sd2 ~ Exponential(1.0)
  t0z ~ Normal(0.0, 0.5)
  dt: data (eltype=Vector{Float64}, n=8)
  stressReport: data (eltype=Vector{Float64}, n=8)
  moodReport: data (eltype=Vector{Float64}, n=8)
  smoked: data (eltype=Vector{Int64}, n=8)
  pred ~ kernel((dti, ys, ym, smk)->begin
        #= brm-docs-example.jl:23 =#
        ys ~ ema_sd(ym, smk, dti, b0, bm, a12, a21, a22, cintm, qd0, qd1, cz, sdm, l31, thr, r1, r2, s0, m0, t0sd1, t0sd2, t0z, nsub, gh)
        #= brm-docs-example.jl:25 =#
        ys
    end, dt, stressReport, moodReport, smoked)
  gh: data (eltype=Int64, n=1)
  nsub: data (eltype=Int64, n=1)
julia
SBBRMI with data keys = [:dt, :gh, :kernel_nsub_pred, :moodReport, :nsub, :smoked, :stressReport]
emitted @slic body:
begin
    b0 ~ normal(0.5, 0.5)
    bm ~ normal(0.4, 0.5)
    a12 ~ normal(-0.25, 0.5)
    a21 ~ normal(-0.3, 0.5)
    a22 ~ normal(-0.6, 0.3)
    cintm ~ normal(0.3, 0.5)
    qd0 ~ normal(-0.2, 0.5)
    qd1 ~ normal(0.3, 0.5)
    cz ~ normal(0.7, 0.5)
    sdm ~ exponential(1.0 ./ 1.0)
    l31 ~ normal(1.2, 0.5)
    thr ~ normal(-1.0, 0.5)
    r1 ~ exponential(1.0 ./ 1.0)
    r2 ~ exponential(1.0 ./ 1.0)
    s0 ~ normal(0.0, 1.0)
    m0 ~ normal(0.5, 1.0)
    t0sd1 ~ exponential(1.0 ./ 1.0)
    t0sd2 ~ exponential(1.0 ./ 1.0)
    t0z ~ normal(0.0, 0.5)
    pred ~ plate(dt, stressReport, moodReport, smoked; outer = (kernel_nsub_pred,)) do dti, ys, ym, smk
            #= brm-docs-example.jl:23 =#
            ys ~ ema_sd(ym, smk, dti, b0, bm, a12, a21, a22, cintm, qd0, qd1, cz, sdm, l31, thr, r1, r2, s0, m0, t0sd1, t0sd2, t0z, nsub, gh)
            #= brm-docs-example.jl:25 =#
            ys
        end
end
stan
functions {
int ragged_end(array[] int ends, int i) {
    return ends[i];
}
int ragged_start(
    array[] int ends,
    int i
) {
    if((i == 1)) {
        return 1;
    } else {
        return (1 + ends[(i - 1)]);
    }
}
real ema_sd_lpdf(
    vector ys,
    vector ym,
    array[] int smoked,
    vector dt,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cintm,
    real qd0,
    real qd1,
    real cz,
    real sdm,
    real l31,
    real thr,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real t0sd1,
    real t0sd2,
    real t0z,
    int nsub,
    int gh
) {
    int T = dims(ys)[1];
    if (dims(ym)[1] != T) reject("ema_sd_lpdf: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(smoked)[1] != T) reject("ema_sd_lpdf: dim mismatch — `smoked` dim 1 (= ", dims(smoked)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("ema_sd_lpdf: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    return sum(
        ema_sd_lpdfs(
            ys,
            ym,
            smoked,
            dt,
            b0,
            bm,
            a12,
            a21,
            a22,
            cintm,
            qd0,
            qd1,
            cz,
            sdm,
            l31,
            thr,
            r1,
            r2,
            ms0,
            mm0,
            t0sd1,
            t0sd2,
            t0z,
            nsub,
            gh
        )
    );
}
vector ema_sd_lpdfs(
    vector ys,
    vector ym,
    array[] int smoked,
    vector dt,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cintm,
    real qd0,
    real qd1,
    real cz,
    real sdm,
    real l31,
    real thr,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real t0sd1,
    real t0sd2,
    real t0z,
    int nsub,
    int gh
) {
    int T = dims(ys)[1];
    if (dims(ym)[1] != T) reject("ema_sd_lpdfs: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(smoked)[1] != T) reject("ema_sd_lpdfs: dim mismatch — `smoked` dim 1 (= ", dims(smoked)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("ema_sd_lpdfs: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `ys` dim 1. `T` sizes: `ys` dim 1 (= ", dims(ys)[1], "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    vector[T] out;
    real ms = ms0;
    real mm = mm0;
    real p11 = (t0sd1 * t0sd1);
    real p12 = (tanh(t0z) * t0sd1 * t0sd2);
    real p22 = (t0sd2 * t0sd2);
    for(t in 1:T) {
        if((t > 1)) {
            real h = (dt[t] / nsub);
            for(st in 1:nsub) {
                if((gh > 0)) {
                    real l11 = sqrt(p11);
                    real l21 = (p12 / l11);
                    real l22 = sqrt(((p22 - (l21 * l21)) + 1.0e-12));
                    real ey1 = 0.0;
                    real ey2 = 0.0;
                    real c11 = 0.0;
                    real c12 = 0.0;
                    real c22 = 0.0;
                    real q11 = 0.0;
                    real q12 = 0.0;
                    for(a in 1:gh) {
                        real xa = ghx(a, gh);
                        for(b in 1:gh) {
                            real xb = ghx(b, gh);
                            real w = (ghw(a, gh) * ghw(b, gh));
                            real xs = (ms + (l11 * xa));
                            real xm = (mm + (l21 * xa) + (l22 * xb));
                            real y1 = (xs + (h * (((-log1p_exp((b0 + (bm * xm)))) * xs) + (a12 * xm))));
                            real y2 = (xm + (h * ((a21 * xs) + (a22 * xm) + cintm)));
                            real sdx = exp((qd0 + (qd1 * xm)));
                            ey1 = (ey1 + (w * y1));
                            ey2 = (ey2 + (w * y2));
                            c11 = (c11 + (w * y1 * y1));
                            c12 = (c12 + (w * y1 * y2));
                            c22 = (c22 + (w * y2 * y2));
                            q11 = (q11 + (w * sdx * sdx));
                            q12 = (q12 + (w * tanh((cz * xs)) * sdx * sdm));
                        }
                    }
                    p11 = ((c11 - (ey1 * ey1)) + (h * q11));
                    p12 = ((c12 - (ey1 * ey2)) + (h * q12));
                    p22 = ((c22 - (ey2 * ey2)) + (h * sdm * sdm));
                    ms = ey1;
                    mm = ey2;
                } else {
                    real sp = log1p_exp((b0 + (bm * mm)));
                    real sig = inv_logit((b0 + (bm * mm)));
                    real nms = (ms + ((((-sp) * ms) + (a12 * mm)) * h));
                    real nmm = (mm + (((a21 * ms) + (a22 * mm) + cintm) * h));
                    real f11 = (1 - (h * sp));
                    real f12 = (h * (a12 - (bm * sig * ms)));
                    real f21 = (h * a21);
                    real f22 = (1 + (h * a22));
                    real sds = exp((qd0 + (qd1 * mm)));
                    real corr = tanh((cz * ms));
                    real qs = (sds * sds * h);
                    real qc = (corr * sds * sdm * h);
                    real qm = (sdm * sdm * h);
                    real fp11 = ((f11 * p11) + (f12 * p12));
                    real fp12 = ((f11 * p12) + (f12 * p22));
                    real fp21 = ((f21 * p11) + (f22 * p12));
                    real fp22 = ((f21 * p12) + (f22 * p22));
                    real np11 = ((fp11 * f11) + (fp12 * f12) + qs);
                    real np12 = ((fp11 * f21) + (fp12 * f22) + qc);
                    real np22 = ((fp21 * f21) + (fp22 * f22) + qm);
                    ms = nms;
                    mm = nmm;
                    p11 = np11;
                    p12 = np12;
                    p22 = np22;
                }
            }
        }
        real v1 = (ys[t] - ms);
        real v2 = (ym[t] - mm);
        real s11 = (p11 + (r1 * r1));
        real s12 = p12;
        real s22 = (p22 + (r2 * r2));
        real det = ((s11 * s22) - (s12 * s12));
        real si11 = (s22 / det);
        real si12 = ((-s12) / det);
        real si22 = (s11 / det);
        real quad = ((v1 * ((si11 * v1) + (si12 * v2))) + (v2 * ((si12 * v1) + (si22 * v2))));
        real lg = (-0.5 * ((2 * l2pi()) + log(det) + quad));
        real k11 = ((p11 * si11) + (p12 * si12));
        real k12 = ((p11 * si12) + (p12 * si22));
        real k21 = ((p12 * si11) + (p22 * si12));
        real k22 = ((p12 * si12) + (p22 * si22));
        ms = (ms + (k11 * v1) + (k12 * v2));
        mm = (mm + (k21 * v1) + (k22 * v2));
        real g11 = (((1 - k11) * p11) - (k12 * p12));
        real g12 = (((1 - k11) * p12) - (k12 * p22));
        real g21 = (((-k21) * p11) + ((1 - k22) * p12));
        real g22 = (((-k21) * p12) + ((1 - k22) * p22));
        p11 = g11;
        p12 = (0.5 * (g12 + g21));
        p22 = g22;
        real etabar = ((l31 * ms) + thr);
        real s2 = ((l31 * l31 * p11) + 1.0e-12);
        real sde = sqrt(s2);
        real z0 = 0.0;
        real z1 = 0.0;
        real z2 = 0.0;
        for(i in 1:5) {
            real xi = ghx(i, 5);
            real wi = ghw(i, 5);
            real eta = (etabar + (sde * xi));
            real pr = (1 - inv_logit(eta));
            if((smoked[t] == 1)) {
                pr = inv_logit(eta);
            }
            z0 = (z0 + (wi * pr));
            z1 = (z1 + (wi * pr * eta));
            z2 = (z2 + (wi * pr * eta * eta));
        }
        real eeta = (z1 / z0);
        real veta = ((z2 / z0) - (eeta * eeta));
        real kb1 = ((p11 * l31) / s2);
        real kb2 = ((p12 * l31) / s2);
        real shr = (s2 - veta);
        ms = (ms + (kb1 * (eeta - etabar)));
        mm = (mm + (kb2 * (eeta - etabar)));
        real n11 = (p11 - (kb1 * kb1 * shr));
        real n12 = (p12 - (kb1 * kb2 * shr));
        real n22 = (p22 - (kb2 * kb2 * shr));
        p11 = n11;
        p12 = n12;
        p22 = n22;
        out[t] = (lg + log(z0));
    }
    return out;
}
real ghx(
    int i,
    int K
) {
    real x = 0.0;
    if((K == 3)) {
        x = ((i - 2) * 1.7320508075688772);
    }
    if((K == 5)) {
        if((i == 1)) {
            x = -2.8569700138728056;
        }
        if((i == 2)) {
            x = -1.355626179974266;
        }
        if((i == 4)) {
            x = 1.355626179974266;
        }
        if((i == 5)) {
            x = 2.8569700138728056;
        }
    }
    return x;
}
real ghw(
    int i,
    int K
) {
    real w = 0.6666666666666666;
    if((K == 3)) {
        if((i != 2)) {
            w = 0.16666666666666666;
        }
    }
    if((K == 5)) {
        w = 0.5333333333333333;
        if(((i == 1) || (i == 5))) {
            w = 0.01125741132772069;
        }
        if(((i == 2) || (i == 4))) {
            w = 0.22207592200561263;
        }
    }
    return w;
}
real l2pi() {
    return 1.8378770664093453;
}
vector ema_sd_vector_rng(
    int anontok__1,
    vector ym,
    array[] int smoked,
    vector dt,
    real b0,
    real bm,
    real a12,
    real a21,
    real a22,
    real cintm,
    real qd0,
    real qd1,
    real cz,
    real sdm,
    real l31,
    real thr,
    real r1,
    real r2,
    real ms0,
    real mm0,
    real t0sd1,
    real t0sd2,
    real t0z,
    int nsub,
    int gh
) {
    int T = anontok__1;
    if (dims(ym)[1] != T) reject("ema_sd_rng: dim mismatch — `ym` dim 1 (= ", dims(ym)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(smoked)[1] != T) reject("ema_sd_rng: dim mismatch — `smoked` dim 1 (= ", dims(smoked)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    if (dims(dt)[1] != T) reject("ema_sd_rng: dim mismatch — `dt` dim 1 (= ", dims(dt)[1], ") does not match `T` (= ", T, "), inferred from `anontok__1` dim 1. `T` sizes: `anontok__1` dim 1 (= ", anontok__1, "), `ym` dim 1 (= ", dims(ym)[1], "), `smoked` dim 1 (= ", dims(smoked)[1], "), `dt` dim 1 (= ", dims(dt)[1], ").");
    vector[T] out;
    real z01 = normal_rng(0.0, 1.0);
    real z02 = normal_rng(0.0, 1.0);
    real r0 = tanh(t0z);
    real s = (ms0 + (t0sd1 * z01));
    real m = (mm0 + (t0sd2 * ((r0 * z01) + (sqrt((1 - (r0 * r0))) * z02))));
    for(t in 1:T) {
        if((t > 1)) {
            real h = (dt[t] / nsub);
            for(st in 1:nsub) {
                real sds = exp((qd0 + (qd1 * m)));
                real corr = tanh((cz * s));
                real zs = normal_rng(0.0, 1.0);
                real zc = normal_rng(0.0, 1.0);
                real z2 = ((corr * zs) + (sqrt((1 - (corr * corr))) * zc));
                real ds = (((((-log1p_exp((b0 + (bm * m)))) * s) + (a12 * m)) * h) + (sds * sqrt(h) * zs));
                real dm = ((((a21 * s) + (a22 * m) + cintm) * h) + (sdm * sqrt(h) * z2));
                s = (s + ds);
                m = (m + dm);
            }
        }
        out[t] = normal_rng(s, r1);
    }
    return out;
}
}
data {
    int kernel_nsub_pred;
    int stressReport_ends_n;
    int stressReport_mem_n;
    tuple(vector[stressReport_mem_n], array[stressReport_ends_n] int) stressReport;
    int dt_ends_n;
    int dt_mem_n;
    tuple(vector[dt_mem_n], array[dt_ends_n] int) dt;
    int moodReport_ends_n;
    int moodReport_mem_n;
    tuple(vector[moodReport_mem_n], array[moodReport_ends_n] int) moodReport;
    int smoked_ends_n;
    int smoked_mem_n;
    tuple(array[smoked_mem_n] int, array[smoked_ends_n] int) smoked;
    int nsub;
    int gh;
}
transformed data {
    array[kernel_nsub_pred] int pred__pl_len_1;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        pred__pl_len_1[plate_i__pl_1] = (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1)));
    }
    array[kernel_nsub_pred] int pred__pl_end_1 = cumulative_sum(pred__pl_len_1);
    vector[sum(pred__pl_len_1)] pred__pl_mem_1;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        pred__pl_mem_1[
            ragged_start(pred__pl_end_1, plate_i__pl_1):ragged_end(pred__pl_end_1, plate_i__pl_1)
        ] = stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ];
    }
}
parameters {
    real b0;
    real bm;
    real a12;
    real a21;
    real a22;
    real cintm;
    real qd0;
    real qd1;
    real cz;
    real<lower=0.0> sdm;
    real l31;
    real thr;
    real<lower=0.0> r1;
    real<lower=0.0> r2;
    real s0;
    real m0;
    real<lower=0.0> t0sd1;
    real<lower=0.0> t0sd2;
    real t0z;
}
transformed parameters {
}
model {
    b0 ~ normal(0.5, 0.5);
    bm ~ normal(0.4, 0.5);
    a12 ~ normal(-0.25, 0.5);
    a21 ~ normal(-0.3, 0.5);
    a22 ~ normal(-0.6, 0.3);
    cintm ~ normal(0.3, 0.5);
    qd0 ~ normal(-0.2, 0.5);
    qd1 ~ normal(0.3, 0.5);
    cz ~ normal(0.7, 0.5);
    sdm ~ exponential((1.0 ./ 1.0));
    l31 ~ normal(1.2, 0.5);
    thr ~ normal(-1.0, 0.5);
    r1 ~ exponential((1.0 ./ 1.0));
    r2 ~ exponential((1.0 ./ 1.0));
    s0 ~ normal(0.0, 1.0);
    m0 ~ normal(0.5, 1.0);
    t0sd1 ~ exponential((1.0 ./ 1.0));
    t0sd2 ~ exponential((1.0 ./ 1.0));
    t0z ~ normal(0.0, 0.5);
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] ~ ema_sd(
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            b0,
            bm,
            a12,
            a21,
            a22,
            cintm,
            qd0,
            qd1,
            cz,
            sdm,
            l31,
            thr,
            r1,
            r2,
            s0,
            m0,
            t0sd1,
            t0sd2,
            t0z,
            nsub,
            gh
        );
    }
}
generated quantities {
    vector[num_elements(stressReport.1)] stressReport_gen;
    vector[num_elements(stressReport.2)] stressReport_likelihood;
    for(plate_i__pl_1 in 1:kernel_nsub_pred) {
        stressReport_gen[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] = ema_sd_vector_rng(
            (1 + (ragged_end(stressReport.2, plate_i__pl_1) - ragged_start(stressReport.2, plate_i__pl_1))),
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            b0,
            bm,
            a12,
            a21,
            a22,
            cintm,
            qd0,
            qd1,
            cz,
            sdm,
            l31,
            thr,
            r1,
            r2,
            s0,
            m0,
            t0sd1,
            t0sd2,
            t0z,
            nsub,
            gh
        );
        stressReport_likelihood[plate_i__pl_1] = ema_sd_lpdf(stressReport.1[
            ragged_start(stressReport.2, plate_i__pl_1):ragged_end(stressReport.2, plate_i__pl_1)
        ] | 
            moodReport.1[ragged_start(moodReport.2, plate_i__pl_1):ragged_end(moodReport.2, plate_i__pl_1)],
            smoked.1[ragged_start(smoked.2, plate_i__pl_1):ragged_end(smoked.2, plate_i__pl_1)],
            dt.1[ragged_start(dt.2, plate_i__pl_1):ragged_end(dt.2, plate_i__pl_1)],
            b0,
            bm,
            a12,
            a21,
            a22,
            cintm,
            qd0,
            qd1,
            cz,
            sdm,
            l31,
            thr,
            r1,
            r2,
            s0,
            m0,
            t0sd1,
            t0sd2,
            t0z,
            nsub,
            gh
        );
    }
}
julia
Turing unsupported for this BRM example

Turing backend: direct execution requires at least one observed likelihood

19 dimensions, whatever the panel size.

The filter's precision is data ​

The last two arguments of ema_sd — nsub, the number of Euler substeps per observation interval, and gh, the predict order (0 first-order, 3 or 5 a Gauss–Hermite rule) — are not literals in the model. They are scalar fields of the data, read inside the cell:

julia
with_filter(d; nsub=8, gh=3) = merge(d, (; nsub, gh))

Every precision is therefore the same compiled Stan model on the same parameter space. That is what makes the precision check further down a loop over log-density evaluations rather than a family of models.

Reading a ctsem specification as a @brm model ​

ctsemhere
DRIFT, CINTthe drift terms inside the @deffun recurrence
DIFFUSION, T0VAR, MANIFESTVARsd / fisher-z cells: the variance is the cell squared, an off-diagonal cell is a correlation through tanh — declared as sds (Exponential) and unconstrained correlations (Normal, then tanh)
TDPREDEFFECT (time-dependent predictors)an impulse added to the state at each observation row, using that row's value
LAMBDA, MANIFESTMEANSloadings and intercepts in the measurement update
binary indicatorsGauss–Hermite integration over the latent predictor inside the filter
T0MEANSinitial state means, global or ~ 1 + (1 | subject)
individual differences (indvarying, the population covariance)random effects on the formula surface; one correlated block is the shared (1 | p | subject)
time-independent predictors (TIpred effects)covariates on the formula surface: b0 ~ 1 + age + treatment + …
integrating over the latent statesthe cell's likelihood is the filter: ys ~ ema_ekf(...)

Fitting: recovery on a simulated panel ​

The state-dependent model is fitted to a panel simulated from known values: 100 subjects × 30 occasions at log-normally spaced intervals (median one time unit), measurement sd 0.3. The posterior is sampled with WarmupHMC's adaptive_warmup_mcmc on the BridgeStan problem that StanBlocks instantiates from the @brm model — 600 draws, 16 substeps, moment-matched predict.

Thirteen of the fourteen 95 % intervals cover the generating value; b0's just misses, and cz sits two posterior sds low (taken up under Sampling variability below). The three state-dependent cells, as functions of the latent state:

The dependence of stress volatility on mood is recovered tightly; the mood-dependence of the recovery rate is the weakest-identified of the three (its two parameters b0, bm trade off against each other).

Is the filter precise enough? ​

The likelihood contains a numerical approximation — the substepped Gaussian filter — and its error biases the posterior by an amount that is unknown a priori. This is the situation that Timonen, Siccha, Bales, Lähdesmäki & Vehtari treat for ODE solvers, and their workflow carries over with the filter in the solver's place:

  1. sample the posterior with a cheap approximation M;

  2. evaluate the same draws under a precise approximation M∗ and Pareto-smooth the importance ratios pM∗(θ∣y)/pM(θ∣y) (PSIS);

  3. a small Pareto k^ certifies M for this posterior, and the reweighted draws are the M∗ posterior at the cost of one likelihood evaluation per draw; a large k^ says: raise the precision and go to 1.

Because the precision is data, the draws of one rung are valid points for every other rung. On the 100 × 30 panel, against the reference M∗ = (32 substeps, 5-node moment-matched predict):

substeps, predictfitsd of log-ratiok^IS-ESS / 600largest shift of a posterior mean vs. certified
2, first-order98 s11.73.2213.95 sd
8, first-order285 s1.920.84340.54 sd
8, moment-matched1715 s2.750.81220.64 sd
16, moment-matched3009 s0.920.142890.21 sd (own mean vs. reweighted)
  • The precision axis that matters is the substep count — the analogue of the solver's step size. Moving from a 3- to a 5-node quadrature changes the log-density by less than 10−6, and moment matching at 8 substeps is no closer to the reference than the first-order predict; halving the step cuts the spread of the log-ratios threefold, as the O(h) error of the Euler scheme predicts.

  • Reweighting 600 draws to the reference takes about 8 minutes; a refit takes 30 to 50. One diagnostic run replaces a precision sweep of MCMC fits.

  • Reweighted estimates are reported only for the rung that passes. An importance-sampling estimate whose k^ exceeds the threshold is not an estimate.

Sampling variability across panels ​

One simulated panel is one draw from the sampling distribution of the estimator, and with fourteen parameters a posterior mean two sds from its generating value is expected somewhere. To separate sampling variability from systematic error, the model is refitted (8 substeps, first-order, about five minutes per panel) on 21 independently simulated panels of the same design, under three generators: two random streams, and a generator eight times finer than the filter's mesh.

generatorpanelsmean czse
Xoshiro stream, 8 Euler–Maruyama steps per interval80.6640.057
Xoshiro stream, 64 steps per interval60.6480.023
LCG stream, 8 steps per interval (includes the panel fitted above)70.5760.053
all210.6300.029

Against the generating value 0.70 and a single-panel posterior sd of 0.115, the estimates scatter around the truth with a shortfall of at most 10 %; the panel used for the recovery figure is the second lowest of the 21. The precision check above rules out the filter's discretisation as the reason for that panel's low value: the certified posterior says 0.476 ± 0.120.

Notes for ctsem users ​

  • Step mesh. The filters here substep every observation interval (nsub), and the generator integrates on a fine mesh. A filter that takes one step per observation interval (ctsem's does unless a maximum time step is set) evaluates state-dependent cells once per interval; data simulated on a finer mesh then come from a different process than the one that filter assumes, and a state-dependent parameter such as cz is the first to show it. When comparing fits across packages, match the meshes.

  • Estimator. The fits on this page are full posteriors by NUTS under weakly informative priors, not maximum-likelihood fits with Hessian-based standard errors.

Reproduce ​

After bootstrapping the repository's test environment (julia --project=test test/setup_env.jl):

sh
# build each model: @brm -> StanBlocks -> stanc -> BridgeStan, finite log-density + gradient
julia --project=test research/ema_ctsem/ema_sampled.jl
julia --project=test research/ema_ctsem/ema_kernel_kalman.jl
julia --project=test research/ema_ctsem/ema_kernel_marginalized.jl
julia --project=test research/ema_ctsem/ema_state_dependent.jl

# fit (nsub gh), the precision ladder, one replication panel
julia --project=test research/ema_ctsem/ema_state_dependent_fit.jl 16 3
julia --project=test research/ema_ctsem/ema_state_dependent_psis.jl OUTDIR
julia --project=test research/ema_ctsem/ema_state_dependent_replicate.jl xo8 1

# tables and figures
julia research/ema_ctsem/ema_state_dependent_summaries.jl OUTDIR
julia --project=research/adaptive_centering/plots research/ema_ctsem/figures.jl

research/ema_ctsem/ema_brm.jl builds the model of section 1 up incrementally, from the pure formula surface to the kernel cell. The measured tables behind the figures are in research/ema_ctsem/results/.

References ​

  • C. C. Driver, J. H. L. Oud, M. C. Voelkle (2017). Continuous Time Structural Equation Modeling with R Package ctsem. Journal of Statistical Software 77(5). doi:10.18637/jss.v077.i05

  • C. C. Driver, M. C. Voelkle (2018). Hierarchical Bayesian Continuous Time Dynamic Modeling. Psychological Methods 23(4), 774–799. doi:10.1037/met0000168

  • J. Timonen, N. Siccha, B. Bales, H. Lähdesmäki, A. Vehtari. An importance sampling approach for reliable and efficient inference in Bayesian ordinary differential equation models. arXiv:2205.09059

  • A. Vehtari, D. Simpson, A. Gelman, Y. Yao, J. Gabry. Pareto Smoothed Importance Sampling. arXiv:1507.02646