Skip to content

fix: keep multidimensional sample_shape in DirectPosterior samples - #2019

Merged
janfb merged 2 commits into
mainfrom
fix/direct-posterior-sample-shape
Oct 5, 2026
Merged

janfb merged 2 commits into
mainfrom
fix/direct-posterior-sample-shape

Conversation

@janfb

@janfb janfb commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do?

DirectPosterior.sample() and .sample_batched() ignore the structure of a multidimensional sample_shape. For sample_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_shape with the existing split_leading_dim helper. DirectPosterior.sample(), DirectPosterior.sample_batched() and NPE_A_Posterior.sample() all use it. NPE_A_Posterior has its own copy of sample(), so it had the same bug.

When return_partial_on_timeout=True returns 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.

Call Before After
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.

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.
@janfb
janfb requested a review from StefanWahl September 30, 2026 10:02
@coderabbitai

coderabbitai Bot commented Sep 30, 2026 •

Copy link
Copy Markdown

Review in Change Stack →

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
  • Configuration used: Organization UI
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: e98d6a73-5ec8-401e-8388-e9469a1b1651
📥 Commits

Reviewing files that changed from the base of the PR and between cc7852c and f00a964.

📒 Files selected for processing (3)
  • sbi/inference/posteriors/direct_posterior.py
  • sbi/utils/torchutils.py
  • tests/posterior_nn_test.py

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


📝 Walkthrough

Walkthrough

Sampling 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.

Changes

Posterior sample shapes

Layer / File(s) Summary
Reshape posterior samples
sbi/inference/posteriors/direct_posterior.py, sbi/inference/posteriors/npe_a_posterior.py, sbi/utils/torchutils.py, tests/posterior_nn_test.py
Direct posterior and NPE-A sampling reshape complete results to the requested sample shape. Direct posterior methods keep partial results flat. split_leading_dim now accepts the Shape type. Tests check the output shapes for single and batched sampling.

Priority: ⬇️ Low

Estimated code review effort: 2 (Simple) | ~8 minutes

Change: Bug fix

Merge Risk: ⚪ Minimal · up to f00a9

Complete samples now preserve the requested shape, and partial timeout results remain flat. No outstanding issue prevents merging.

Security Architecture Review

Security architecture risk: ⚪ Minimal · up to f00a9

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
No architecture-level concerns identified.

Security review details

Security Blast Radius

  • inferred — The verified change is confined to in-process posterior result dimensions and consumers of that contract. The inspected changes introduce no additional tenant, service, credential, or persistent-store access.

Trust Boundaries and Controls

  • inferred — Caller-provided sample_shape affects the layout of samples already produced through the existing sampling path. The new helper neither invokes another sampler nor changes observation validation, prior-support policy, or sampling authority.
🚥 Pre-merge checks | ✅ 5
✅ Passed checks (5 passed)
Check name Status Explanation
Docstring Coverage ✅ Passed Docstring coverage is 90.00% which is sufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 10 functions across 4 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.
Title check ✅ Passed The title clearly and concisely describes the main change: preserving multidimensional sample_shape in DirectPosterior samples.
Description check ✅ Passed The description explains the problem, changes, behavior differences, tests, and follow-up context. It omits the template checklist and does not link an issue, but it is otherwise substantially complet…
✨ Finishing Touches
📝 Generate docstrings
  • Commit to this branch
  • Create a new PR
🧪 Generate unit tests (beta)
  • Commit to this branch
  • Create a new PR

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 Sep 30, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 87.50000% with 1 line in your changes missing coverage. Please review.
✅ Project coverage is 89.32%. Comparing base (bcb50b9) to head (f00a964).
⚠️ Report is 21 commits behind head on main.
✅ All tests successful. No failed tests found.

Files with missing lines Patch % Lines
sbi/inference/posteriors/direct_posterior.py 80.00% 1 Missing ⚠️
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     
Flag Coverage Δ
fast 85.03% <87.50%> (?)

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

Files with missing lines Coverage Δ
sbi/inference/posteriors/npe_a_posterior.py 80.23% <100.00%> (+23.25%) ⬆️
sbi/utils/torchutils.py 80.25% <100.00%> (+2.15%) ⬆️
sbi/inference/posteriors/direct_posterior.py 82.41% <80.00%> (+0.20%) ⬆️

... and 53 files with indirect coverage changes

@dgedon dgedon left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread sbi/inference/posteriors/direct_posterior.py
Comment thread sbi/inference/posteriors/direct_posterior.py Outdated
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.
@janfb

janfb commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor Author

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

@janfb
janfb merged commit b8afbd0 into main Oct 5, 2026
21 checks passed
@janfb
janfb deleted the fix/direct-posterior-sample-shape branch October 5, 2026 12:19
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.

2 participants