mirror of
https://github.com/NixOS/nixpkgs.git
synced 2026-08-25 17:55:21 +00:00
Diff: https://github.com/pyro-ppl/funsor/compare/0.4.7...0.4.8 Changelog: https://github.com/pyro-ppl/funsor/releases/tag/0.4.8
53 lines
2.4 KiB
Diff
53 lines
2.4 KiB
Diff
diff --git a/funsor/distribution.py b/funsor/distribution.py
|
|
index 5a48ecb..cfb2d2e 100644
|
|
--- a/funsor/distribution.py
|
|
+++ b/funsor/distribution.py
|
|
@@ -309,7 +309,10 @@ class Distribution(Funsor, metaclass=DistributionMeta):
|
|
@classmethod
|
|
@functools.lru_cache(maxsize=5000)
|
|
def _infer_param_domain(cls, name, raw_shape):
|
|
- support = cls.dist_class.arg_constraints.get(name, None)
|
|
+ constraints = getattr(cls.dist_class, "arg_constraints", {})
|
|
+ if isinstance(constraints, property):
|
|
+ constraints = {}
|
|
+ support = constraints.get(name, None)
|
|
# XXX: if the backend does not have the same definition of constraints, we should
|
|
# define backend-specific distributions and overide these `infer_value_domain`,
|
|
# `infer_param_domain` methods.
|
|
@@ -330,9 +333,8 @@ class Distribution(Funsor, metaclass=DistributionMeta):
|
|
# for discrete multivariate distributions in Pyro
|
|
elif support_name == "Real":
|
|
if name == "logits" and (
|
|
- "probs" in cls.dist_class.arg_constraints
|
|
- and type(cls.dist_class.arg_constraints["probs"]).__name__.lstrip("_")
|
|
- == "Simplex"
|
|
+ "probs" in constraints
|
|
+ and type(constraints.get("probs")).__name__.lstrip("_") == "Simplex"
|
|
):
|
|
output = Reals[raw_shape[-1 - event_dim :]]
|
|
else:
|
|
@@ -377,10 +379,13 @@ def make_dist(
|
|
backend_dist_class, param_names=(), generate_eager=True, generate_to_funsor=True
|
|
):
|
|
if not param_names:
|
|
+ constraints = getattr(backend_dist_class, "arg_constraints", {})
|
|
+ if isinstance(constraints, property):
|
|
+ constraints = {}
|
|
param_names = tuple(
|
|
name
|
|
for name in inspect.getfullargspec(backend_dist_class.__init__)[0][1:]
|
|
- if name in backend_dist_class.arg_constraints
|
|
+ if name in constraints
|
|
)
|
|
|
|
@makefun.with_signature(
|
|
@@ -608,7 +613,7 @@ class CoerceDistributionToFunsor:
|
|
def __call__(self, cls, args, kwargs):
|
|
# Check whether distribution class takes any tensor inputs.
|
|
arg_constraints = getattr(cls, "arg_constraints", None)
|
|
- if not arg_constraints:
|
|
+ if not arg_constraints or isinstance(arg_constraints, property):
|
|
return
|
|
|
|
# Check whether any tensor inputs are actually funsors.
|