Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions numpyro/distributions/discrete.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,8 +68,8 @@ def _to_probs_multinom(logits: ArrayLike) -> ArrayLike:


def _to_logits_multinom(probs: ArrayLike) -> ArrayLike:
minval = jnp.finfo(jnp.result_type(probs)).min
return jnp.clip(jnp.log(probs), minval)
safe_probs = jnp.where(probs > 0, probs, 1.0)
return jnp.where(probs > 0, jnp.log(safe_probs), -jnp.inf)


class BernoulliProbs(Distribution):
Expand Down
7 changes: 4 additions & 3 deletions numpyro/distributions/mixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -220,9 +220,10 @@ def __init__(
*,
validate_args: Optional[bool] = None,
):
assert isinstance(
component_distribution.support, constraints.ParameterFreeConstraint
), (
base_support = component_distribution.support
while isinstance(base_support, constraints._IndependentConstraint):
base_support = base_support.base_constraint
assert isinstance(base_support, constraints.ParameterFreeConstraint), (
f"Invalid component distribution: {type(component_distribution).__name__}. "
"The mixture components must have a support that does not depend on their parameters "
f"(expected ParameterFreeConstraint, but found {component_distribution.support})."
Expand Down
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ dev = [
"matplotlib",
"optax>=0.0.6",
"pandas>=3.0.3",
"pre-commit>=4.6.1",
"pylab-sdk", # jaxns dependency
"pyro-api>=0.1.2",
"pytest>=9.1.1",
Expand Down
4 changes: 1 addition & 3 deletions test/test_distributions_mixture.py
Original file line number Diff line number Diff line change
Expand Up @@ -188,9 +188,7 @@ def _test_mixture(mixing_distribution, component_distribution):
)
def test_mixture_rejects_parameter_dependent_components(component_dist):
mixing_dist = dist.Categorical(probs=np.array([0.5, 0.5]))
with pytest.raises(
AssertionError, match="expected ParameterFreeConstraint, but found "
):
with pytest.raises(AssertionError, match="ParameterFreeConstraint, but found "):
dist.MixtureSameFamily(mixing_dist, component_dist)


Expand Down
Loading