Skip to content

Run build_smc_logfns and SMC prior particles consistently in unconstrained space - #266

Merged
junpenglao merged 1 commit into
mainfrom
smc-logfns-space
Oct 2, 2026
Merged

junpenglao merged 1 commit into
mainfrom
smc-logfns-space

Conversation

@junpenglao

Copy link
Copy Markdown
Member

Summary

The documented contract is that SMC positions are unconstrained: postprocess_fn maps unconstrained → constrained, and init_position is param_info.z. Three pieces broke that contract:

  • logprior_fn evaluated position as constrained values through log_density, with no transform and no Jacobian.
  • loglikelihood_fn = −potential_fn − logprior_fn therefore mixed the two conventions, since numpyro's potential_fn takes unconstrained values.
  • build_prior_sample_fn's Predictive fallback returned constrained particles.

The fix:

  • logprior_fn(z) = −blocked_potential_fn(z), where blocked_potential_fn is the potential of the prior-only (observed sites blocked) model, built by its own initialize_model. It uses the same unconstrained convention and Jacobian as the joint potential_fn, so the Jacobians cancel and loglikelihood_fn(z) is exactly log p(data | constrain(z)).
  • The prior particles are mapped through each site's inverse bijector, so they come out unconstrained.
  • The docstrings are corrected.
  • The joint initialize_model still receives rng_key unchanged. init_position is bit-identical to main, checked on 2 models × 3 seeds.

Impact on results: none for existing end-to-end SMC runs. logprior_fn + loglikelihood_fn ≡ −potential_fn held by construction, so the λ=1 target was always correct. The old constrained particles were also self-consistent with the old logprior_fn at λ=0. As a result, adaptive tempered SMC with an RWM kernel recovered the NUTS posterior mean (within about 0.005) on a toy model with both old and new code. The bug affected consumers that use logprior_fn, loglikelihood_fn or the particles on their own, and kernels that need genuinely unconstrained positions. Of the committed SMC recipes, gmm_25 and neals_funnel use the analytic-sampler path and logistic_synthetic has only real-valued latents, so none of them hit the broken branch.

Found during the adversarial review of #265 (trax local Issue#1266).

Test plan

  • New tests/numpyro/test_smc_logfns_space.py, a toy model with sigma ~ HalfNormal(2) and y ~ Normal(0, sigma).
    • Tests that fail on main:
      • logprior_fn and loglikelihood_fn against an independent analytic reference at 7 unconstrained z. On main, loglikelihood_fn is off at all 7, e.g. z=2: −18.92 vs −14.60.
      • Prior particles include negative values, and their postprocessed mean recovers E[HalfNormal(2)] = 1.60. On main it is 14.11.
    • Structural invariants, which pass on both: logprior + loglik == −potential, and the blocked model's latent sites match the joint model's.
    • Of the 6 new tests, 4 fail on main.
  • Existing SMC tests: test_emit_smc_script.py, test_smc_generated_lifecycle.py, test_generated_smc.py, tests/smc/, test_api_pins_smc.py, test_registry.py (162 passed together with the new file).
  • pytest tests -m fast -n 2: 2089 passed, 5 skipped.
  • pre-commit clean.
  • CI green.

🤖 Generated with Claude Code

…d space mix

Finding: build_smc_logfns's logprior_fn used log_density (direct value
substitution, no transform -- wants constrained-space position) while
loglikelihood_fn = -potential_fn(position) - logprior_fn(position) used
potential_fn from initialize_model (explicitly unconstrained-space per its
own docstring: "transform these unconstrained parameters to the values
belong to the supports"). Same "position" argument, two incompatible
conventions. build_prior_sample_fn's Predictive fallback compounded this:
it returned constrained-space particles (Predictive's own native output)
with no inverse-transform step, despite init_position/postprocess_fn
(unconstrained -> constrained) already establishing unconstrained as the
module's actual contract.

Verified with a toy model (sigma ~ HalfNormal(2.0), y ~ Normal(0, sigma) --
one positive-support latent, so constrained and unconstrained genuinely
differ) against an analytic reference built from NumPyro's own
biject_to(positive) bijector, independent of the code under test: on main,
loglikelihood_fn mismatched the analytic likelihood at 7/7 sampled
unconstrained points (e.g. z=2.0: -18.92 vs analytic -14.60), and
build_prior_sample_fn's particles, pushed through postprocess_fn, gave a
mean of 14.13 against HalfNormal(2.0)'s true mean of 1.60 (particles were
already-constrained, so postprocess_fn's own transform double-applied).

Investigated but NOT asserted by a new test (end-to-end behavior, for the
record): because loglikelihood_fn is *defined* as
-potential_fn(w) - logprior_fn(w), logprior_fn(w) + loglikelihood_fn(w) ==
-potential_fn(w) is a tautology of that formula regardless of what
logprior_fn itself computes -- so the full-tempering (lambda=1) SMC target
was always mathematically exactly -potential_fn(w) in main's code too, and
main's own build_prior_sample_fn particles were self-consistent with
main's own (buggy) logprior_fn's labeled convention at lambda=0. Run
end-to-end (blackjax.adaptive_tempered_smc, RWM inner kernel, same pattern
the production lifecycle template uses) on the toy model above, both main
and this fix land within ~0.005 of a NUTS reference mean for sigma. The
break is in logprior_fn/loglikelihood_fn/particles used individually or by
a consumer that needs a genuinely unconstrained position (e.g. an inner
HMC-family kernel's leapfrog dynamics, or any code that inspects the prior
alone) -- not in this one self-contained round trip.

Fix: logprior_fn is now the negated potential energy of the prior-only
(observed-sites-blocked) model, built via the model's own initialize_model
call -- the same unconstrained-space + Jacobian convention the joint
potential_fn uses, so the two compose correctly and the Jacobian terms
cancel exactly in loglikelihood_fn. build_prior_sample_fn's Predictive
fallback now derives each site's inverse bijector once (from a single
seeded trace -- biject_to(support) depends only on the declared support,
not the realized value) and applies it batched to the whole particle set
before returning. Both docstrings corrected to state the unconstrained
convention explicitly.

Verification: new tests in tests/numpyro/test_smc_logfns_space.py pin the
toy-model analytic equivalence and the unconstrained-particle invariant;
confirmed failing on main (git stash of the fix) before the fix and
passing after. Existing SMC tests (test_emit_smc_script.py,
test_smc_generated_lifecycle.py, test_generated_smc.py, tests/smc/,
test_api_pins_smc.py, test_registry.py) all still pass -- none of the
three models with committed SMC recipes (gmm_25, neals_funnel,
logistic_synthetic) exercised the broken path in a way their own tests
could catch: the first two use the analytic_sampler fast path (bypasses
build_prior_sample_fn's buggy branch entirely) and logistic_synthetic's
one latent has real (unconstrained-identity) support, so the space bug was
invisible there. `pytest tests -m fast -n 2` (worked around an unrelated,
pre-existing pytest-benchmark/pytest-xdist incompatibility on this
worktree's current lock with `-p no:benchmark`; already fixed by the
separate in-flight uv.lock upgrade) -- 2089 passed, 5 skipped (6 more than
baseline, matching the new tests). Pre-commit clean.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@junpenglao
junpenglao merged commit e04e65f into main Oct 2, 2026
4 checks passed
@junpenglao
junpenglao deleted the smc-logfns-space branch October 2, 2026 13:52
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant