Repository navigation
Run build_smc_logfns and SMC prior particles consistently in unconstrained space - #266
Merged
Merged
Conversation
…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>
2 tasks done
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
The documented contract is that SMC positions are unconstrained:
postprocess_fnmaps unconstrained → constrained, andinit_positionisparam_info.z. Three pieces broke that contract:logprior_fnevaluatedpositionas constrained values throughlog_density, with no transform and no Jacobian.loglikelihood_fn = −potential_fn − logprior_fntherefore mixed the two conventions, since numpyro'spotential_fntakes unconstrained values.build_prior_sample_fn'sPredictivefallback returned constrained particles.The fix:
logprior_fn(z) = −blocked_potential_fn(z), whereblocked_potential_fnis the potential of the prior-only (observed sites blocked) model, built by its owninitialize_model. It uses the same unconstrained convention and Jacobian as the jointpotential_fn, so the Jacobians cancel andloglikelihood_fn(z)is exactly log p(data | constrain(z)).initialize_modelstill receivesrng_keyunchanged.init_positionis bit-identical tomain, checked on 2 models × 3 seeds.Impact on results: none for existing end-to-end SMC runs.
logprior_fn + loglikelihood_fn ≡ −potential_fnheld by construction, so the λ=1 target was always correct. The old constrained particles were also self-consistent with the oldlogprior_fnat λ=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 uselogprior_fn,loglikelihood_fnor the particles on their own, and kernels that need genuinely unconstrained positions. Of the committed SMC recipes,gmm_25andneals_funneluse the analytic-sampler path andlogistic_synthetichas 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
tests/numpyro/test_smc_logfns_space.py, a toy model withsigma ~ HalfNormal(2)andy ~ Normal(0, sigma).main:logprior_fnandloglikelihood_fnagainst an independent analytic reference at 7 unconstrainedz. Onmain,loglikelihood_fnis off at all 7, e.g. z=2: −18.92 vs −14.60.mainit is 14.11.logprior + loglik == −potential, and the blocked model's latent sites match the joint model's.main.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.🤖 Generated with Claude Code