Skip to content

Conditionally non-centre latent terminal drift (#124) - #125

Open
ErikRingen wants to merge 13 commits into
mainfrom
124-conditional-noncentred-terminal-drift
Open

ErikRingen wants to merge 13 commits into
mainfrom
124-conditional-noncentred-terminal-drift

Conversation

@ErikRingen

Copy link
Copy Markdown
Collaborator

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 on Q_sigma and A, 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 pointwise log_lik, applied in the other direction.

Change

  • Scope. The new parameterisation applies when Gaussian variables are present, taxa are not duplicated, and some terminal drift is latent (non-Gaussian variables or missing Gaussian values). This is decided by an internal use_conditional_ncp(). Models with duplicated taxa and estimate_residual = FALSE keep the centred form, because one latent value there is shared by several observations.
  • Stan.
    • Transformed data computes a per-observation permutation that puts observed Gaussian variables first. No standata changes are needed.
    • ncp_terminal_drift() maps innovations to the realised drift using the permuted Cholesky factor. It reuses L_VCV_tips when the permutation is the identity.
    • The model block evaluates the same centred multivariate normal density on the realised drift, and adds the log Jacobian outside the log_sum_exp over trees. The change of variables is therefore exact for multiPhylo models too.
    • Generated quantities compute log_lik from the realised drift.
  • JAX. The nutpie backend mirrors the Stan change, using a masked triangular solve so shapes stay static. coev_make_model_config() passes a conditional_ncp flag to the JAX model.
  • The Stan code for every other configuration is byte-identical to main. I checked 56 configurations: all response types, repeated observations, duplicated taxa, measurement error, exact and approximate GPs, and multiPhylo, each with log_lik and prior_only on and off.

Model equivalence

New tests are in tests/testthat/test-conditional_ncp.R and helper-conditional_ncp.R:

  • The helper regenerates the centred code by mocking 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 that
    • log_prob(new) == log_prob(centred) + log|J| to within 1e-8 relative tolerance, and
    • pointwise log_lik is identical to within 1e-10.
  • Configurations tested:
    • normal + Bernoulli;
    • missing data in all variables, with a non-Gaussian variable first;
    • restricted effects_mat without correlated drift;
    • ordinal + normal + Poisson;
    • measurement error;
    • multiPhylo;
    • normal-only with missing values;
    • exact GP.
  • Two configurations run by default; the rest need COEVOLVE_EXTENDED_TESTS=true, following the existing Stan/JAX suite.
  • Mutation tests: removing the Jacobian fails the log-density check, and removing the GQ transform fails the log_lik check.
  • compare_stan_jax_logprob() gains a log_lik option. New Stan/JAX tests compare log density and pointwise log_lik for 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.
  • The full default suite passes with NOT_CRAN=true, as does the full extended Stan/JAX suite.
  • testthat (>= 3.1.7) in Suggests, for with_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:

main branch
Q_sigma[P] bulk / tail ESS 147 / 251 1,918 / 1,469
Q_sigma[P] R-hat 1.023 1.003
E-BFMI 0.53–0.74 0.89–1.12
Posterior means agree (|z| < 0.8 against MCSE)

40-tip full model (cross-effects, correlated drift, missing data), adapt_delta = 0.99, 1,000 + 2,000 iterations:

main branch
Q_sigma[z] bulk / tail ESS 177 / 85 5,209 / 3,239
R-hat 1.037 1.000
Divergences 6 1
E-BFMI 0.18–0.59 0.92–1.04
Posterior means (12 parameters) max |z| = 1.75
elpd_loo −60.5 (SE 7.3) −60.0 (SE 7.4)

Wall 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_logit support 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

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>
@ScottClaessens

Copy link
Copy Markdown
Owner

Thanks for getting this started @ErikRingen. It looks like we're having some issues regenerating the test fixtures and with running the tests in test-logp_stan_jax.R. I'll look into this more tomorrow.

@ScottClaessens

Copy link
Copy Markdown
Owner

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:

Variables: x = normal 
           y = bernoulli_logit 
     Data: d (Number of observations: 50)
Phylogeny: tree (Number of trees: 1)
    Draws: 4 chains, each with iter = 1000; warmup = 1000; thin = 1
           total post-warmup draws = 4000

Autoregressive selection effects:
  Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
x    -1.62      0.81 -3.29 -0.20 1.00     2704     1559
y    -1.08      0.68 -2.63 -0.08 1.00     2667     1770

Cross selection effects:
      Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
x ⟶ y    -0.01      0.62 -1.30  1.27 1.00     1571     1878
y ⟶ x     0.01      1.00 -1.94  1.92 1.00     6554     2775

Drift parameters:
         Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
sd(x)       10.63      0.45  9.78 11.53 1.00     6842     3113
sd(y)        0.59      0.47  0.04  1.82 1.17       21       55
cor(x,y)     0.08      0.35 -0.62  0.68 1.06       48      613

Continuous time intercept parameters:
  Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
x     0.03      0.96 -1.84  1.92 1.00     7689     2929
y    -0.22      0.77 -1.67  1.36 1.00     3872     2593

Warning message:
Parts of the model have not converged (some Rhats are > 1.05). Be careful when analysing the results! We recommend running more iterations and/or setting stronger priors.

This PR branch:

Variables: x = normal 
           y = bernoulli_logit 
     Data: d (Number of observations: 50)
Phylogeny: tree (Number of trees: 1)
    Draws: 4 chains, each with iter = 1000; warmup = 1000; thin = 1
           total post-warmup draws = 4000

Autoregressive selection effects:
  Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
x    -1.61      0.82 -3.31 -0.19 1.00     2655     1575
y    -1.09      0.68 -2.60 -0.06 1.00     2709     1413

Cross selection effects:
      Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
x ⟶ y    -0.01      0.59 -1.25  1.18 1.00     2635     2644
y ⟶ x     0.00      1.03 -1.99  1.99 1.00     6455     2965

Drift parameters:
         Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
sd(x)       10.63      0.46  9.77 11.56 1.00     8803     2743
sd(y)        0.60      0.48  0.02  1.79 1.00     3374     2169
cor(x,y)     0.01      0.34 -0.63  0.64 1.00    10214     2971

Continuous time intercept parameters:
  Estimate Est.Error  2.5% 97.5% Rhat Bulk_ESS Tail_ESS
x     0.00      0.98 -1.88  1.97 1.00     7278     3159
y    -0.24      0.76 -1.70  1.27 1.00     4772     3073

@ScottClaessens

Copy link
Copy Markdown
Owner

Not entirely sure why the R CMD checks are failing to regenerate fixtures. I was able to run tests/testthat/fixtures/coevfit_examples.R locally with this branch.

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 estimate_correlated_drift = FALSE. As a reproducible example, the authority model compiles and runs normally when using default settings, but not when disabling correlated drift.

coev_fit(
  data = authority$data,
  variables = list(
    religious_authority = "ordered_logistic",
    political_authority = "ordered_logistic"
  ),
  id = "language",
  tree = authority$phylogeny,
  estimate_correlated_drift = FALSE
)
Error: An error occured during compilation! See the message above for more information.

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.

@ScottClaessens

Copy link
Copy Markdown
Owner

One potentially informative error that crops up repeatedly in the logs is:

error: no matching function for call to 'diag_matrix'

But I'm not sure why this occurs when correlated drift is disabled but not when it's enabled, as diag_matrix is used in the Stan code for both versions.

@ScottClaessens

Copy link
Copy Markdown
Owner

Strangely, it seems to be that CmdStan 2.40.0 doesn't like the use of diag_matrix(Q_sigma^2) in the transformed parameters block. I've changed this to diag_matrix(square(Q_sigma)), which seems to fix the problem.

@ScottClaessens ScottClaessens added the run-extended Run the extended (logp Stan/JAX) test suite in CI label Sep 28, 2026
@codecov

codecov Bot commented Sep 28, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.

📢 Thoughts on this report? Let us know!

@ScottClaessens

Copy link
Copy Markdown
Owner

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.

@yuhangxoox

Copy link
Copy Markdown

Thanks both — I also tested the package-level implementation on our original BALANCED190_A exact-D0 diagnostic using the public coev_fit() interface, without any package or generated-Stan modifications.

For clarity, this validation used PR head 20e32cc09a5116e06ddb110f8561abb26ba7d84d, so it predates the most recent commits to this PR.

The conditional NCP was activated as expected. For Q_sigma[P], the original centred exact-D0 fit had R-hat = 1.0688, bulk ESS = 49.0 and tail ESS = 31.1; under the PR implementation these improved to R-hat = 1.0074, bulk ESS = 826.6 and tail ESS = 780.9. The minimum chain E-BFMI also improved from 0.242 to 0.902.

The full D0 still did not converge satisfactorily: the remaining problems were concentrated mainly in the H–P cross-effect / drift-correlation block (P → H, H → P, and the H–P drift correlation), with maximum R-hat = 1.506, some divergences, and substantial max-treedepth saturation. This is very similar to what we saw with our earlier custom conditional-NCP diagnostic.

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.

@ScottClaessens

Copy link
Copy Markdown
Owner

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 estimate_correlated_drift = FALSE and see if that makes a difference.

@ScottClaessens

Copy link
Copy Markdown
Owner

The error message for the remaining failing test is:

Error: Error: Exception: cholesky_corr_free: x is not a valid unit vector. The sum of the squares of the elements should be 1, but is 1 (in '/tmp/RtmpJ18sgQ/model-1c343c94e805.stan', line 144, column 2 to column 37)

Which suggests to me that this is perhaps a rounding issue?

@yuhangxoox

Copy link
Copy Markdown

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 estimate_correlated_drift = FALSE.

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. Q_sigma[P] remained well behaved (R-hat 1.006).

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. eta_anc[P] and b[P] show the same 2+2 shift in chain means, although their marginal distributions overlap substantially.

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.

@ErikRingen

Copy link
Copy Markdown
Collaborator Author

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 estimate_correlated_drift = FALSE.

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. Q_sigma[P] remained well behaved (R-hat 1.006).

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. eta_anc[P] and b[P] show the same 2+2 shift in chain means, although their marginal distributions overlap substantially.

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?

@yuhangxoox

Copy link
Copy Markdown

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!

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

run-extended Run the extended (logp Stan/JAX) test suite in CI

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Conditional non-centering improves terminal-drift sampling in mixed Gaussian–Bernoulli models

3 participants