Skip to content

fix: apply the NPE-A proposal correction through the estimator - #2045

Open
janfb wants to merge 8 commits into
mainfrom
refactor/npe-a-corrected-estimator
Open

janfb wants to merge 8 commits into
mainfrom
refactor/npe-a-corrected-estimator

Conversation

@janfb

@janfb janfb commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

Summary

From round 2 on, NPE-A corrects the network's MoG for the proposal. Since PR #2035, two DirectPosterior hooks apply this correction. Code that uses the estimator directly misses it: map(), potential() and ConditionedMDN use the raw network, and .to() leaves the correction on the old device.

Now NPE_A_Posterior wraps its MDN in a ProposalCorrectedMDN from round 2 on. Its sample() and log_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.

 NPE_A_Posterior, round 2 on
-  posterior_estimator: MDN, plus hooks that correct sample() and log_prob()
+  posterior_estimator: ProposalCorrectedMDN(MDN, proposal_mog, prior_mog)
 ProposalCorrectedMDN.sample / log_prob
+  MDN.sample_from_mog / log_prob_from_mog with the corrected MoG
 ProposalCorrectedMDN.to()
+  also moves and casts the proposal and prior MoGs
 extract_and_transform_mog, used by ConditionedMDN
+  corrected MoG of a ProposalCorrectedMDN
 DirectPosterior
-  _sample_estimator / _log_prob_estimator hooks
 _correct_for_proposal and its helpers
-  sbi/inference/trainers/npe/npe_a.py
+  sbi/neural_nets/estimators/mog.py              # pure move

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.

round-2 NPE-A main this PR
potential() at θ = (0, 0) −3.30, the raw network −2.86, equal to log_prob()
map() (−0.23, 0.32), the raw mode (−0.54, 0.55), the corrected mode
after .to("mps") sample(), log_prob() and map() fail with a device mismatch all methods run
ConditionedMDN conditions the raw MoG conditions the corrected MoG

sample(), log_prob() and the batched methods give the same values as on main.

Both new tests fail on main and pass with this PR:

test_npe_a_map_potential_and_conditional_apply_proposal_correction
  potential(θ) == log_prob(θ, norm_posterior=False)
  map() == closed-form corrected mode
  ConditionedMDN.log_prob(θ₀) − log_prob((θ₀, 1)) is the same for all θ₀
test_npe_a_corrected_posterior_to_device   (gpu marked)
  after .to(gpu): log_prob() and potential() match the CPU values; sample() and map() run

Merge Danger

Door: two-way. ProposalCorrectedMDN is not exported, and the removed hooks were private.

Blast Radius: multi-round NPE-A.

From round 2 on, posterior.posterior_estimator is a ProposalCorrectedMDN, not a MixtureDensityEstimator. 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.

ConditionedMDN gives wrong results when a free dimension comes after a fixed one. That bug is in MoG.condition and also affects other MDNs; PR #2039 fixes it.

janfb added 8 commits October 6, 2026 10:00
…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.
@janfb
janfb requested a review from dgedon October 6, 2026 11:36
@coderabbitai

coderabbitai Bot commented Oct 6, 2026 •

Copy link
Copy Markdown

Review in Change Stack →

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: Organization UI
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: 1c61a7b5-ce5e-4f72-9b2c-8f5ac4e4b8eb
📥 Commits

Reviewing files that changed from the base of the PR and between 29af137 and 640c72f.

📒 Files selected for processing (10)
  • sbi/analysis/conditional_density.py
  • sbi/inference/posteriors/direct_posterior.py
  • sbi/inference/posteriors/npe_a_posterior.py
  • sbi/inference/trainers/npe/npe_a.py
  • sbi/neural_nets/estimators/mixture_density_estimator.py
  • sbi/neural_nets/estimators/mog.py
  • sbi/utils/conditional_density_utils.py
  • tests/inference_on_device_test.py
  • tests/leakage_correction_test.py
  • tests/posterior_nn_test.py
💤 Files with no reviewable changes (1)
  • sbi/inference/trainers/npe/npe_a.py

Included review availability: This review used your included allowance. Your plan provides up to 4 included reviews per hour; 3 remain after this review.


📝 Walkthrough

Walkthrough

SNPE-A proposal correction now runs through a ProposalCorrectedMDN estimator. The estimator provides corrected MoG scoring and sampling, and NPE-A posterior and conditional-density code accepts the corrected estimator.

Changes

SNPE-A proposal correction

Layer / File(s) Summary
MoG proposal-correction calculations
sbi/neural_nets/estimators/mog.py, sbi/inference/trainers/npe/npe_a.py
The correction calculations and helper operations move from the NPE-A trainer module to the MoG module. The correction forms proposal–density component pairs and computes corrected precisions, means, and logits.
Corrected MDN estimator
sbi/neural_nets/estimators/mixture_density_estimator.py
The MDN gains explicit-MoG scoring and sampling methods. ProposalCorrectedMDN stores proposal and optional prior MoGs, computes corrected MoGs, and uses them for scoring and sampling. Its loss raises NotImplementedError.
NPE-A posterior integration
sbi/inference/posteriors/npe_a_posterior.py, sbi/inference/posteriors/direct_posterior.py
NPE-A passes the corrected estimator to DirectPosterior when a proposal MoG is supplied. DirectPosterior calls estimator sampling and log-probability methods directly; NPE-A MoG parameter retrieval delegates to the estimator.
Conditional-density support and validation
sbi/analysis/conditional_density.py, sbi/utils/conditional_density_utils.py, tests/posterior_nn_test.py, tests/inference_on_device_test.py, tests/leakage_correction_test.py
Conditional-density code accepts ProposalCorrectedMDN and extracts its corrected MoG. Tests cover corrected posterior operations, device movement, and leakage-correction sampler spying.

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
Loading

Suggested reviewers: bharath0153

Merge Risk: ⚪ Minimal · up to 640c7

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 build_posterior(), a limitation the PR description states. No merge-blocking risk was found.

🚥 Pre-merge checks | ✅ 5
✅ Passed checks (5 passed)
Check name Status Explanation
Title check ✅ Passed The title clearly summarizes the main change: applying NPE-A proposal correction through the estimator.
Description check ✅ Passed The description explains the change, its motivation, test evidence, and known compatibility limitations. It omits the template’s issue-link and checklist sections, but is otherwise detailed and comple…
Docstring Coverage ✅ Passed Docstring coverage is 88.57% which is sufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 35 functions across 9 files.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
✨ Finishing Touches
📝 Generate docstrings
  • Commit to this branch
  • Create a new PR
🧪 Generate unit tests (beta)
  • Commit to this branch
  • Create a new PR
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

@codecov

codecov Bot commented Oct 6, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 90.19608% with 10 lines in your changes missing coverage. Please review.
✅ Project coverage is 89.86%. Comparing base (29af137) to head (640c72f).
✅ All tests successful. No failed tests found.

Files with missing lines Patch % Lines
...eural_nets/estimators/mixture_density_estimator.py 76.47% 8 Missing ⚠️
sbi/neural_nets/estimators/mog.py 96.42% 2 Missing ⚠️
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     
Flag Coverage Δ
fast 85.45% <90.19%> (?)

Flags with carried forward coverage won't be shown. Click here to find out more.

Files with missing lines Coverage Δ
sbi/analysis/conditional_density.py 96.66% <ø> (+3.33%) ⬆️
sbi/inference/posteriors/direct_posterior.py 84.87% <100.00%> (-0.50%) ⬇️
sbi/inference/posteriors/npe_a_posterior.py 100.00% <100.00%> (ø)
sbi/inference/trainers/npe/npe_a.py 73.72% <ø> (-7.31%) ⬇️
sbi/utils/conditional_density_utils.py 86.86% <100.00%> (+0.29%) ⬆️
sbi/neural_nets/estimators/mog.py 94.61% <96.42%> (+0.60%) ⬆️
...eural_nets/estimators/mixture_density_estimator.py 89.16% <76.47%> (-2.89%) ⬇️

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant