Fix MoG conditioning for non-contiguous dimensions - #2040
vansh-oberoi wants to merge 2 commits into
Conversation
|
Navigate logical layers of code changes, visualize relationships, and explore their blast radius. Note Currently processing new changes in this PR. This may take a few minutes, please wait... ⚙️ Run configuration
📒 Files selected for processing (1)
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (2)
Included review availability: This review used your included allowance. Your plan provides up to 4 included reviews per hour; 3 remain after this review. 📝 WalkthroughWalkthrough
ChangesMoG conditioning
Priority: ➖ Normal Estimated code review effort: 3 (Moderate) | ~20 minutes Change: Bug fix · Severity of issue fixed: Medium Merge Risk: ⚪ Minimal · up to This change corrects MoG conditioning when a free dimension follows a fixed one, and adds an analytic regression test. No merge-blocking risk is evident from the supplied review. 🚥 Pre-merge checks | ✅ 5✅ Passed checks (5 passed)
✨ Finishing Touches🧪 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 |
What does this PR do?
Fixes #2038.
The conditioning of a mixture of Gaussians was incorrect when a free dimension came after a fixed dimension.
The issue was caused by constructing the precision submatrices by slicing the precision factor before computing the full precision matrix. This can give incorrect conditional precisions for non-contiguous dimensions.
This change computes the full precision matrix first, extracts the required blocks, and uses the Schur complement when calculating the marginal likelihood used for the mixture weights.
I also added a regression test using an analytically known Gaussian conditional.
Tests
uv run pytest tests/mog_test.py -quv run pytest tests/sbiutils_test.py -quv run pytest tests/posterior_nn_test.py -quv run ruff check sbi/neural_nets/estimators/mog.py tests/mog_test.pyAll passed.