Conditionally non-centre latent terminal drift (#124) - #125
ErikRingen wants to merge 13 commits into
Conversation
In models with Gaussian variables and no repeated observations, the latent terminal drift (non-Gaussian variables and missing Gaussian values) was sampled on its realised scale, creating a funnel with the drift parameters. It is now written as a standard normal innovation conditional on the observed Gaussian residuals of the same taxon. The model block evaluates the same density on the realised drift and adds the change-of-variables Jacobian outside the mixture over trees, so the model, posterior and pointwise log_lik are unchanged. The JAX backend mirrors the change. Stan code for all other configurations is byte-identical to main. Tests compare the new code against the centred code (regenerated by mocking use_conditional_ncp()) across missing data, measurement error, multiPhylo, GPs, variable ordering and response types, and compare Stan and JAX log density and pointwise log_lik. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
|
Thanks for getting this started @ErikRingen. It looks like we're having some issues regenerating the test fixtures and with running the tests in |
|
Just confirming myself that this PR improves sampling for a Gaussian-Bernoulli model. # set seed
set.seed(123)
# simulate data
n <- 50
tree <- ape::rcoal(n)
d <- data.frame(
id = tree$tip.label,
x = rnorm(n),
y = as.integer(sample(0:1, size = n, replace = TRUE))
)
# fit model
fit <- coev_fit(
data = d,
variables = list(
x = "normal",
y = "bernoulli_logit"
),
id = "id",
tree = tree,
parallel_chains = 4,
seed = 123
)
# print summary
summary(fit)Main branch: This PR branch: |
|
Not entirely sure why the R CMD checks are failing to regenerate fixtures. I was able to run GitHub Actions is installing CmdStan version 2.40.0, whereas locally I'm using 2.38.0. I updated to the most recent version and re-ran the test fixtures. The following model failed to run: # tests/testthat/fixtures/coevfit_examples.R, line 197
coevfit_example_09 <-
coev_fit(
data = d,
variables = list(
x = "bernoulli_logit",
y = "bernoulli_logit"
),
id = "id",
tree = tree,
estimate_correlated_drift = FALSE,
chains = chains,
iter_warmup = warmup,
iter_sampling = iter,
adapt_delta = 0.99,
seed = 12345
)This issue seems to be triggered when setting coev_fit(
data = authority$data,
variables = list(
religious_authority = "ordered_logistic",
political_authority = "ordered_logistic"
),
id = "language",
tree = authority$phylogeny,
estimate_correlated_drift = FALSE
)The backlogs are long and I won't copy them here. I don't think this is necessarily due to the changes implemented in this branch, as the same error occurs when running the above example on the main branch with CmdStan 2.40.0. I will dig into this further. |
|
One potentially informative error that crops up repeatedly in the logs is: But I'm not sure why this occurs when correlated drift is disabled but not when it's enabled, as |
…ed parameters Stan block
|
Strangely, it seems to be that CmdStan 2.40.0 doesn't like the use of |
Codecov Report✅ All modified and coverable lines are covered by tests. 📢 Thoughts on this report? Let us know! |
|
Just pushed some changes that will hopefully deal with most of the failures. The one remaining test failure is the following extended test: # tests/testthat/test-logp_stan_jax.R, line 400
test_that("logp agrees WITH likelihood: repeated measures", {
skip_if_not(run_extended_tests)
expect_logp_agreement(
data = repeated$data,
variables = list(x = "normal", y = "normal"),
id = "species",
tree = repeated$phylogeny,
prior_only = FALSE,
grad_tol = 1e-2
)
})@ErikRingen could you look into this for me? It's probably just a tolerance issue. |
|
Thanks both — I also tested the package-level implementation on our original BALANCED190_A exact-D0 diagnostic using the public For clarity, this validation used PR head The conditional NCP was activated as expected. For The full D0 still did not converge satisfactorily: the remaining problems were concentrated mainly in the H–P cross-effect / drift-correlation block ( So, at least on our original real-data diagnostic, the package implementation seems to cleanly resolves the specific Bernoulli terminal-drift mixing problem we reported in #124, while the remaining full-D0 geometry appears to be a separate issue. Thanks again for implementing this and for all the work on the tests. |
|
Thanks for the info @yuhangxoox. I think as Erik alluded to before, the issues you are having probably stem from the strong correlation between H and P. If the model includes correlated drift, it is likely finding it difficult to disentangle whether the correlation arises from the cross-selection effects or the correlated drift. I would try turning off the correlated drift with |
|
The error message for the remaining failing test is: Which suggests to me that this is perhaps a rounding issue? |
|
Thanks Scott — I tested your suggestion on the same BALANCED190_A exact-D0 diagnostic. We kept the same PR commit, tree, data, scaling, effects matrix, off-diagonal prior, seed and Stan controls; the sole model change was Removing correlated drift removed the observed HMC trajectory pathologies in this run: divergences went from 34/2000 to 0, and max-treedepth hits from 1537/2000 (76.85%) to 0. Minimum chain E-BFMI remained good at 0.868. However, the reciprocal H–P effects did not mix better. P→H changed from R-hat 1.506 / bulk ESS 7.43 to 1.733 / 6.17, and H→P from 1.205 / 14.16 to 1.734 / 6.20. I also inspected the four saved CmdStan chains without refitting. The separation is very structured: chains 1–2 stayed near P→H ≈ +7.08 and H→P ≈ −2.09, whereas chains 3–4 stayed near P→H ≈ −7.21 and H→P ≈ +2.34 throughout all 500 saved sampling iterations. The paired draws occupy two clearly separated chain-specific regions, with no observed transitions between them. So disabling correlated drift removes the divergences and treedepth saturation, but it does not resolve the remaining H↔P mixing failure. In these saved chains, the pattern looks more like distinct posterior regions than a single slowly explored ridge, although the finite nonconverged sample cannot establish that the full posterior is truly multimodal. Just wanted to pass this along as an additional real-data diagnostic in case it is useful for interpreting the remaining geometry. Hope this helps. |
That's interesting. I would interpret the sign reversal as: the model cannot determine directionality in this sample/tree. It sees they are obviously correlated at the tips, but there's no sufficient information to leverage in the tree. Out of curiosity, did you ever try re-running with the whole tree? |
|
Not yet with the new parameterisation — I’ll rerun the current PR #125 no-correlated-drift version on the whole tree and see whether the sign reversal persists. Thanks! |
Closes #124.
Problem
When a model has Gaussian variables and no repeated observations,
coev_make_stancode()samples latent terminal drift on its realised scale. Latent drift here means the drift of non-Gaussian variables and of missing Gaussian values. Its scale depends onQ_sigmaandA, which creates a funnel between the drift parameters and the tip-level latent drift. Models without Gaussian variables already use a non-centred terminal drift, which is why, as @yuhangxoox found, a Bernoulli trait mixes well alone but poorly once a Gaussian trait is added. This happens even when cross-effects and correlated drift are disabled and the two processes are independent.The centred form was used because observed Gaussian residuals are fixed by the data, so the full tip vector cannot be written as
L * z(see #55). This PR splits the tip density into p(observed residuals) × p(latent drift | observed residuals) and non-centres only the conditional part. This is the same conditional-normal algebra the generated quantities already use for the pointwiselog_lik, applied in the other direction.Change
use_conditional_ncp(). Models with duplicated taxa andestimate_residual = FALSEkeep the centred form, because one latent value there is shared by several observations.standatachanges are needed.ncp_terminal_drift()maps innovations to the realised drift using the permuted Cholesky factor. It reusesL_VCV_tipswhen the permutation is the identity.log_sum_expover trees. The change of variables is therefore exact for multiPhylo models too.log_likfrom the realised drift.coev_make_model_config()passes aconditional_ncpflag to the JAX model.log_likandprior_onlyon and off.Model equivalence
New tests are in
tests/testthat/test-conditional_ncp.Randhelper-conditional_ncp.R:use_conditional_ncp(). It then maps random points from the new parameterisation to the centred one using an independent R implementation of the transform. It checks thatlog_prob(new) == log_prob(centred) + log|J|to within 1e-8 relative tolerance, andlog_likis identical to within 1e-10.effects_matwithout correlated drift;COEVOLVE_EXTENDED_TESTS=true, following the existing Stan/JAX suite.log_likcheck.compare_stan_jax_logprob()gains alog_likoption. New Stan/JAX tests compare log density and pointwiselog_likfor the new path (missing data, measurement error, multiPhylo). The ~1e-5 Stan/JAX differences match main on the same data, because the JAX matrix exponential is approximate.NOT_CRAN=true, as does the full extended Stan/JAX suite.testthat (>= 3.1.7)in Suggests, forwith_mocked_bindings().Sampling
All fits used
coev_fit()on main vs this branch with identical data and seeds, 4 chains.190-tip birth–death tree, E (normal) + P (Bernoulli), no cross-effects or correlated drift, 500 + 500 iterations:
Q_sigma[P]bulk / tail ESSQ_sigma[P]R-hat40-tip full model (cross-effects, correlated drift, missing data),
adapt_delta = 0.99, 1,000 + 2,000 iterations:Q_sigma[z]bulk / tail ESSelpd_looWall time is similar in both, because the step size is set by the continuous traits. As noted in #124, larger models with strongly correlated continuous traits can still have geometry problems that this change does not address.
Note for the categorical branch
categorical_logitsupport is not on main yet. When that branch is rebased,use_conditional_ncp()and the permutation should use latent indices (latent_var_id) rather than variable indices.🤖 Generated with Claude Code