Repository navigation
Conversation
…sityEstimator log_prob() and sample() now delegate to these helpers, which evaluate or sample a given MoG with the estimator's input transform. No behavior change.
Pure move of _correct_for_proposal() and its helpers from the NPE-A trainer to the estimators package, so an estimator can use it. NPE_A_Posterior now imports it at module level instead of lazily from the trainer.
From round 2 on, NPE_A_Posterior wraps its estimator in ProposalCorrectedMDN, whose sample() and log_prob() use the corrected MoG. The potential is built from that estimator, so map() and potential() are corrected too. Before, they used the raw network, and map() returned its mode. NPE_A_Posterior no longer overrides the DirectPosterior hooks.
Only NPE_A_Posterior overrode them, and it now corrects through its estimator. DirectPosterior calls posterior_estimator.sample() and log_prob() directly again. _alternative_sampling_method stays for NPE-A's low-acceptance hint.
ProposalCorrectedMDN now moves its proposal and prior MoGs in _apply(), so NPE_A_Posterior.to() moves the whole correction. Before, round-2 NPE-A failed with a device mismatch in sample(), log_prob() and map() after .to(), and potential() ran uncorrected.
extract_and_transform_mog() takes the corrected MoG from a ProposalCorrectedMDN and the z-score from the MDN it wraps. Before, ConditionedMDN of a round-2 NPE-A estimator conditioned the uncorrected MoG.
The wrapper's log_prob() and sample() have the same shapes as the wrapped MDN, so their docstrings point there. The ConditionedMDN check moves into the map()/potential() test, which builds the same posterior.
ProposalCorrectedMDN._apply() now applies the module's fn to the proposal and prior MoGs, so .double() casts them like the network. Before, it moved them by device only, and they stayed float32.
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (10)
💤 Files with no reviewable changes (1)
Included review availability: This review used your included allowance. Your plan provides up to 4 included reviews per hour; 3 remain after this review. 📝 WalkthroughWalkthroughSNPE-A proposal correction now runs through a ChangesSNPE-A proposal correction
Priority: ➖ Normal Estimated code review effort: 3 (Moderate) | ~25 minutes Change: Bug fix Sequence Diagram(s)sequenceDiagram
participant NPE_A_Posterior
participant DirectPosterior
participant ProposalCorrectedMDN
participant MixtureDensityEstimator
NPE_A_Posterior->>ProposalCorrectedMDN: Wrap estimator with proposal and prior MoGs
DirectPosterior->>ProposalCorrectedMDN: Call log_prob or sample
ProposalCorrectedMDN->>MixtureDensityEstimator: Pass corrected MoG to explicit-MoG method
Suggested reviewers: Merge Risk: ⚪ Minimal · up to From round 2 onward, multi-round NPE-A now applies the proposal correction through the estimator, so sampling, log-probability, potential, MAP and conditional densities use the corrected mixture. Later rounds still receive the corrected proposal. With a corrected NPE-A proposal, NPE-C uses its atomic-loss fallback instead of failing. Previously pickled round-2 posteriors load without the correction until they are rebuilt with 🚥 Pre-merge checks | ✅ 5✅ Passed checks (5 passed)
✨ Finishing Touches📝 Generate docstrings
🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
Codecov Report❌ Patch coverage is
Additional details and impacted files@@ Coverage Diff @@
## main #2045 +/- ##
==========================================
- Coverage 89.89% 89.86% -0.04%
==========================================
Files 142 142
Lines 14646 14656 +10
==========================================
+ Hits 13166 13170 +4
- Misses 1480 1486 +6
Flags with carried forward coverage won't be shown. Click here to find out more.
|
Summary
From round 2 on, NPE-A corrects the network's MoG for the proposal. Since PR #2035, two
DirectPosteriorhooks apply this correction. Code that uses the estimator directly misses it:map(),potential()andConditionedMDNuse the raw network, and.to()leaves the correction on the old device.Now
NPE_A_Posteriorwraps its MDN in aProposalCorrectedMDNfrom round 2 on. Itssample()andlog_prob()use the corrected MoG. So everything that uses the posterior's estimator is corrected, and the hooks are removed again. This is Daniel's idea 2 from the PR #2035 review.The wrapper evaluates and samples the corrected MoG with two new MDN methods. So it reuses the MDN's input transform instead of a copy of it. This is idea 3.
This matters for everyone who runs NPE-A for more than one round.
Evidence
Round-2 NPE-A with an untrained one-component MDN and a Gaussian proposal. The corrected posterior is then a Gaussian with a known mode.
potential()at θ = (0, 0)log_prob()map().to("mps")sample(),log_prob()andmap()fail with a device mismatchConditionedMDNsample(),log_prob()and the batched methods give the same values as on main.Both new tests fail on main and pass with this PR:
Merge Danger
Door: two-way.
ProposalCorrectedMDNis not exported, and the removed hooks were private.Blast Radius: multi-round NPE-A.
From round 2 on,
posterior.posterior_estimatoris aProposalCorrectedMDN, not aMixtureDensityEstimator. So NPE-C with a round-2 NPE-A proposal now uses its atomic loss, which is correct for any proposal. Before, it used its cheaper MoG loss with the raw MoG of that proposal, which is the wrong proposal density.Round-2 NPE-A posteriors that were pickled before this PR load without the correction. NPE-A is rarely used, so I did not add a migration. The workaround is to call
build_posterior()again.ConditionedMDNgives wrong results when a free dimension comes after a fixed one. That bug is inMoG.conditionand also affects other MDNs; PR #2039 fixes it.