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
with Gaussian measurement error for continuous indicators and a logistic link for binary ones. The drift
Three layers, one model
Every model on this page divides the work the same way:
| layer | holds | written as |
|---|---|---|
| formula surface | the population model: covariates and (correlated) random effects on subject-level parameters | b0 ~ 1 + age + treatment + (1 | p | subject) |
kernel(...) cell | one subject: that subject's series and that subject's parameter values | pred ~ kernel(dt, stressReport, …, b0, …) do … end |
@deffun | the recurrence — a loop with carried state, emitted as a Stan function | an 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:
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
endThe 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).
Hierarchical EMA model, latent states sampledfunction 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
endBRMI:
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)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
endfunctions {
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)
];
}
}Turing unsupported for this BRM example
Turing backend: direct execution requires at least one observed likelihoodOn 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(...).
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
endLinear-Gaussian model, states Kalman-marginalized per subjectfunction 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
endBRMI:
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)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
endfunctions {
// 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
);
}
}Turing unsupported for this BRM example
Turing backend: direct execution requires at least one observed likelihoodNo 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.
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
endHierarchical EMA model, states EKF-marginalized per subjectfunction 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
endBRMI:
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)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
endfunctions {
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
);
}
}Turing unsupported for this BRM example
Turing backend: direct execution requires at least one observed likelihoodThis 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,
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.
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
endAll 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.
State-dependent drift and diffusion, states marginalized per subjectfunction 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
endBRMI:
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)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
endfunctions {
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
);
}
}Turing unsupported for this BRM example
Turing backend: direct execution requires at least one observed likelihood19 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:
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
| ctsem | here |
|---|---|
DRIFT, CINT | the drift terms inside the @deffun recurrence |
DIFFUSION, T0VAR, MANIFESTVAR | sd / 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, MANIFESTMEANS | loadings and intercepts in the measurement update |
| binary indicators | Gauss–Hermite integration over the latent predictor inside the filter |
T0MEANS | initial 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 states | the 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:
sample the posterior with a cheap approximation
; evaluate the same draws under a precise approximation
and Pareto-smooth the importance ratios (PSIS); a small Pareto
certifies for this posterior, and the reweighted draws are the posterior at the cost of one likelihood evaluation per draw; a large 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
| substeps, predict | fit | sd of log-ratio | IS-ESS / 600 | largest shift of a posterior mean vs. certified | |
|---|---|---|---|---|---|
| 2, first-order | 98 s | 11.7 | 3.22 | 1 | 3.95 sd |
| 8, first-order | 285 s | 1.92 | 0.84 | 34 | 0.54 sd |
| 8, moment-matched | 1715 s | 2.75 | 0.81 | 22 | 0.64 sd |
| 16, moment-matched | 3009 s | 0.92 | 0.14 | 289 | 0.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
, 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 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
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.
| generator | panels | mean cz | se |
|---|---|---|---|
| Xoshiro stream, 8 Euler–Maruyama steps per interval | 8 | 0.664 | 0.057 |
| Xoshiro stream, 64 steps per interval | 6 | 0.648 | 0.023 |
| LCG stream, 8 steps per interval (includes the panel fitted above) | 7 | 0.576 | 0.053 |
| all | 21 | 0.630 | 0.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 asczis 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):
# 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.jlresearch/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



