fix: keep multidimensional sample_shape in DirectPosterior samples - #2019
Conversation
DirectPosterior.sample and sample_batched returned the samples flat, for example (6, D) instead of (2, 3, D) for sample_shape (2, 3). The other posteriors keep the full sample_shape, and get_posterior_samples_on_batch expects it since #2012. NPE_A_Posterior.sample had a copy of the same return line. All three now reshape through one helper. A partial result on timeout holds fewer samples than sample_shape and stays flat. As a side effect, the default sample_shape () now gives (D,) instead of (1, D), as for the other posteriors.
|
Navigate logical layers of code changes, visualize relationships, and explore their blast radius. No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (3)
Included review availability: This review used your included allowance. Your plan provides up to 4 included reviews per hour; 2 remain after this review. 📝 WalkthroughWalkthroughSampling methods now reshape complete results to match the requested sample shape. Direct posterior methods keep partial results flat. Regression tests check single and batched sampling shapes. ChangesPosterior sample shapes
Priority: ⬇️ Low Estimated code review effort: 2 (Simple) | ~8 minutes Change: Bug fix Merge Risk: ⚪ Minimal · up to Complete samples now preserve the requested shape, and partial timeout results remain flat. No outstanding issue prevents merging. Security Architecture ReviewSecurity architecture risk: ⚪ Minimal · up to The change affects returned tensor dimensions, not sampling authority or security controls. Complete results preserve the requested sample dimensions; partial timeout results retain their existing flat form. Callers relying on the old empty-shape result will need to accommodate the removed leading dimension. Retained concerns Security review detailsSecurity Blast Radius
Trust Boundaries and Controls
🚥 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 #2019 +/- ##
==========================================
- Coverage 89.44% 89.32% -0.12%
==========================================
Files 142 142
Lines 14509 18695 +4186
==========================================
+ Hits 12977 16699 +3722
- Misses 1532 1996 +464
Flags with carried forward coverage won't be shown. Click here to find out more.
|
dgedon
left a comment
There was a problem hiding this comment.
Thanks, the fix works. I checked sample((2, 3)) and sample_batched((2, 3), ...) for NPE-C and NPE-A, and the diagnostics from #2012 now pass on both the batched and the fallback path. I left two inline comments: one on the default-shape change, one small reuse.
A follow-up for #2035: this bug had to be fixed in two places because NPE_A_Posterior.sample() is a near-copy of DirectPosterior.sample(); the only real difference is which sampler it calls. main now has the _sample_estimator hook (#2033), and #2035 already routes sample_batched through it. If DirectPosterior.sample() used the hook too, NPE_A_Posterior.sample() could be deleted. I think that fits better in #2035 than here.
The reshape helper now calls split_leading_dim, and its type hint widens from List[int] to Shape. The test also checks that the empty default sample shape gives (D,) for sample and (B, D) for sample_batched.
|
Thanks for the review, both points done and committed on top, PR description updated. Merging now and we are doing the follow-ups in #2035 |
What does this PR do?
DirectPosterior.sample()and.sample_batched()ignore the structure of a multidimensionalsample_shape. Forsample_shape=(2, 3)they return shape(6, D)and(6, B, D), not(2, 3, D)and(2, 3, B, D). The docstring promises the reshape, and all other posteriors (MCMC, rejection, importance, vector field, ensemble) already do it.PR #2012 fixed
get_posterior_samples_on_batch()for multidimensional sample shapes. With NPE posteriors, the diagnostics still failed at its shape check. This was the open CodeRabbit comment on PR #2012. The fix helps everyone who passes a multidimensional sample shape to an NPE posterior, for example through SBC or TARP.The fix in this PR reshapes the samples to
sample_shapewith the existingsplit_leading_dimhelper.DirectPosterior.sample(),DirectPosterior.sample_batched()andNPE_A_Posterior.sample()all use it.NPE_A_Posteriorhas its own copy ofsample(), so it had the same bug.When
return_partial_on_timeout=Truereturns fewer samples than requested, the samples stay flat, as before.Breaking change for the default
sample_shape=(). This is the same behaviour as for the other posteriors. No internal caller uses the default shape.sample()(1, D)(D,)sample_batched((), x)(1, B, D)(B, D)Does this close any issues?
Follow-up to PR #2012.
Anything else we should know?
The new test checks both methods for NPE-A and NPE-C, with
sample_shape=(2, 3)and with the default shape. The full fast test suite passes locally.Created with the help of Claude Code, and reviewed locally with GPT Sol 6.