Skip to content

Fix compute_marginals and sample_posterior for plates of size 1 - #3484

Open
MohammadHijjawi97 wants to merge 1 commit into
pyro-ppl:devfrom
MohammadHijjawi97:fix-compute-marginals-size-1-plate
Open

MohammadHijjawi97 wants to merge 1 commit into
pyro-ppl:devfrom
MohammadHijjawi97:fix-compute-marginals-size-1-plate

Conversation

@MohammadHijjawi97

Copy link
Copy Markdown

Fixes #3102

Proposed changes

TraceEnum_ELBO.compute_marginals() and TraceEnum_ELBO.sample_posterior() raise KeyError when an enumerated site is inside a plate of size 1. Packed tensors drop size-1 dims, so the plate's dim never appears in any factor; contract_to_tensor() was called with only a cache, so its default LogRing had an empty dim_to_size and Ring.broadcast() could not look up the plate size when broadcasting the result back to the site's ordinal.

This passes a LogRing whose dim_to_size is filled from the site's plate frames, in both _compute_marginals and BackwardSampleMessenger.

Tests

  • New test_marginals_plate_3102[1,2,3] in tests/infer/test_enum.py checks the marginals against the analytic posterior and runs sample_posterior; size=1 fails with KeyError: 'a' before this change. The script from the issue now prints tensor([0.3000, 0.7000]).
  • pytest tests/infer/test_enum.py (678 passed) and pytest tests/infer/test_discrete.py.
  • ruff check / ruff format --check / mypy on the changed files.

TraceEnum_ELBO.compute_marginals() and sample_posterior() raised a
KeyError when an enumerated site lived in a plate of size 1. Packed
tensors drop size-1 dims, so the default LogRing had no size to
broadcast the result back to the site's plate. Pass a ring that knows
the site's plate sizes.

Fixes pyro-ppl#3102

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.

TraceEnum_ELBO.compute_marginals: KeyError for batch size = 1

1 participant