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
10 changes: 0 additions & 10 deletions .flake8

This file was deleted.

7 changes: 2 additions & 5 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -10,18 +10,15 @@ docs: FORCE
$(MAKE) -C docs html

lint: FORCE
flake8
black --check .
isort --check .
ruff check --fix .
python scripts/update_headers.py --check
python test/test_import.py

license: FORCE
python scripts/update_headers.py

format: license FORCE
black .
isort .
ruff format .

test: lint FORCE
ifeq (${FUNSOR_BACKEND}, torch)
Expand Down
2 changes: 1 addition & 1 deletion docs/source/conf.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@

if "READTHEDOCS" not in os.environ:
# if developing locally, use funsor.__version__ as version
from funsor import __version__ # noqaE402
from funsor import __version__ # noqa: E402

version = __version__

Expand Down
2 changes: 1 addition & 1 deletion examples/forward_backward.py
Original file line number Diff line number Diff line change
Expand Up @@ -162,7 +162,7 @@ def main(args):
actual = adj + trans - Z
assert_close(expected, actual.align(tuple(expected.inputs)), rtol=1e-4)
print("")
print(f"Marginal term: p(x[{t}], x[{t-1}] | Y)")
print(f"Marginal term: p(x[{t}], x[{t - 1}] | Y)")
print("Forward-backward algorithm:\n", expected.data)
print("Differentiating forward algorithm:\n", actual.data)
t += 1
Expand Down
8 changes: 6 additions & 2 deletions examples/mixed_hmm/experiment.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,12 +166,16 @@ def closure():
re_str = "g" + (
"n"
if args["group"] is None
else "d" if args["group"] == "discrete" else "c"
else "d"
if args["group"] == "discrete"
else "c"
)
re_str += "i" + (
"n"
if args["individual"] is None
else "d" if args["individual"] == "discrete" else "c"
else "d"
if args["individual"] == "discrete"
else "c"
)
results_filename = "expt_{}_{}_{}.json".format(
args["dataset"], re_str, str(uuid.uuid4().hex)[0:5]
Expand Down
4 changes: 3 additions & 1 deletion examples/talbot.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,9 @@

@make_funsor
def InverseLaplace(
F: Has[{"s"}], t: Funsor, s: Bound # noqa: F821
F: Has[{"s"}],
t: Funsor,
s: Bound,
) -> Fresh[lambda F: F]:
"""
Inverse Laplace transform of function F(s).
Expand Down
4 changes: 2 additions & 2 deletions funsor/adjoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,8 +152,8 @@ def _fail_default(*args):
adjoint_ops = KeyedRegistry(default=_fail_default)
if instrument.DEBUG:
adjoint_ops_register = adjoint_ops.register
adjoint_ops.register = lambda *args: lambda fn: adjoint_ops_register(*args)(
instrument.debug_logged(fn)
adjoint_ops.register = lambda *args: (
lambda fn: adjoint_ops_register(*args)(instrument.debug_logged(fn))
)


Expand Down
8 changes: 2 additions & 6 deletions funsor/distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -819,9 +819,7 @@ def eager_multinomial(total_count, probs, value):
else:
total_count = Tensor(ops.expand(total_count, shape[:-1]), inputs)
backend_dist = import_module(BACKEND_TO_DISTRIBUTIONS_BACKEND[get_backend()])
return backend_dist.Multinomial.eager_log_prob(
total_count, probs, value
) # noqa: F821
return backend_dist.Multinomial.eager_log_prob(total_count, probs, value) # noqa: F821


def eager_categorical_funsor(probs, value):
Expand All @@ -841,9 +839,7 @@ def eager_delta_tensor(v, log_density, value):
event_dim = len(v.output.shape)
inputs, (v, log_density, value) = align_tensors(v, log_density, value)
backend_dist = import_module(BACKEND_TO_DISTRIBUTIONS_BACKEND[get_backend()])
data = backend_dist.Delta.dist_class(v, log_density, event_dim).log_prob(
value
) # noqa: F821
data = backend_dist.Delta.dist_class(v, log_density, event_dim).log_prob(value) # noqa: F821
return Tensor(data, inputs)


Expand Down
2 changes: 1 addition & 1 deletion funsor/domains.py
Original file line number Diff line number Diff line change
Expand Up @@ -337,7 +337,7 @@ def _find_domain_getitem(op, lhs_domain, rhs_domain):
elif isinstance(lhs_domain, ProductDomain):
# XXX should this return a Union?
raise NotImplementedError(
"Cannot statically infer domain from: " f"{lhs_domain}[{rhs_domain}]"
f"Cannot statically infer domain from: {lhs_domain}[{rhs_domain}]"
)


Expand Down
4 changes: 2 additions & 2 deletions funsor/interpretations.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,8 +121,8 @@ def __init__(self, name="dispatched"):
self.registry = registry = KeyedRegistry(default=lambda *args: None)

if instrument.DEBUG or instrument.PROFILE:
self.register = lambda *args: lambda fn: registry.register(*args)(
instrument.debug_logged(fn)
self.register = lambda *args: (
lambda fn: registry.register(*args)(instrument.debug_logged(fn))
)
else:
self.register = registry.register
Expand Down
52 changes: 15 additions & 37 deletions funsor/jax/distributions.py
Original file line number Diff line number Diff line change
Expand Up @@ -274,25 +274,19 @@ def categorical_to_funsor(numpyro_dist, output=None, dim_to_name=None):
new_pyro_dist = _NumPyroWrapper_Binomial(
total_count=numpyro_dist.total_count, probs=numpyro_dist.probs
)
return backenddist_to_funsor(
Binomial, new_pyro_dist, output, dim_to_name
) # noqa: F821
return backenddist_to_funsor(Binomial, new_pyro_dist, output, dim_to_name) # noqa: F821


@to_funsor.register(dist.CategoricalProbs)
def categorical_to_funsor(numpyro_dist, output=None, dim_to_name=None):
new_pyro_dist = _NumPyroWrapper_Categorical(probs=numpyro_dist.probs)
return backenddist_to_funsor(
Categorical, new_pyro_dist, output, dim_to_name
) # noqa: F821
return backenddist_to_funsor(Categorical, new_pyro_dist, output, dim_to_name) # noqa: F821


@to_funsor.register(dist.GeometricProbs)
def categorical_to_funsor(numpyro_dist, output=None, dim_to_name=None):
new_pyro_dist = _NumPyroWrapper_Geometric(probs=numpyro_dist.probs)
return backenddist_to_funsor(
Geometric, new_pyro_dist, output, dim_to_name
) # noqa: F821
return backenddist_to_funsor(Geometric, new_pyro_dist, output, dim_to_name) # noqa: F821


@to_funsor.register(dist.MultinomialProbs)
Expand All @@ -301,9 +295,7 @@ def categorical_to_funsor(numpyro_dist, output=None, dim_to_name=None):
new_pyro_dist = _NumPyroWrapper_Multinomial(
total_count=numpyro_dist.total_count, probs=numpyro_dist.probs
)
return backenddist_to_funsor(
Multinomial, new_pyro_dist, output, dim_to_name
) # noqa: F821
return backenddist_to_funsor(Multinomial, new_pyro_dist, output, dim_to_name) # noqa: F821


@to_funsor.register(dist.Delta) # Delta **distribution**
Expand All @@ -323,22 +315,16 @@ def deltadist_to_funsor(pyro_dist, output=None, dim_to_name=None):
]


eager.register(Beta, Funsor, Funsor, Funsor)(eager_beta) # noqa: F821)
eager.register(Beta, Funsor, Funsor, Funsor)(eager_beta) # noqa: F821
eager.register(Binomial, Funsor, Funsor, Funsor)(eager_binomial) # noqa: F821
eager.register(Multinomial, Tensor, Tensor, Tensor)(eager_multinomial) # noqa: F821)
eager.register(Categorical, Funsor, Tensor)(eager_categorical_funsor) # noqa: F821)
eager.register(Categorical, Tensor, Variable)(eager_categorical_tensor) # noqa: F821)
eager.register(Multinomial, Tensor, Tensor, Tensor)(eager_multinomial) # noqa: F821
eager.register(Categorical, Funsor, Tensor)(eager_categorical_funsor) # noqa: F821
eager.register(Categorical, Tensor, Variable)(eager_categorical_tensor) # noqa: F821
eager.register(Delta, Tensor, Tensor, Tensor)(eager_delta_tensor) # noqa: F821
eager.register(Delta, Funsor, Funsor, Variable)(
eager_delta_funsor_variable
) # noqa: F821
eager.register(Delta, Variable, Funsor, Variable)(
eager_delta_funsor_variable
) # noqa: F821
eager.register(Delta, Funsor, Funsor, Variable)(eager_delta_funsor_variable) # noqa: F821
eager.register(Delta, Variable, Funsor, Variable)(eager_delta_funsor_variable) # noqa: F821
eager.register(Delta, Variable, Funsor, Funsor)(eager_delta_funsor_funsor) # noqa: F821
eager.register(Delta, Variable, Variable, Variable)(
eager_delta_variable_variable
) # noqa: F821
eager.register(Delta, Variable, Variable, Variable)(eager_delta_variable_variable) # noqa: F821
eager.register(Normal, Funsor, Tensor, Funsor)(eager_normal) # noqa: F821
eager.register(MultivariateNormal, Funsor, Tensor, Funsor)(eager_mvn) # noqa: F821
eager.register(
Expand All @@ -356,25 +342,17 @@ def deltadist_to_funsor(pyro_dist, output=None, dim_to_name=None):
)( # noqa: F821
eager_dirichlet_multinomial
)
eager.register(
Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Gamma
)( # noqa: F821
eager.register(Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Gamma)( # noqa: F821
eager_gamma_gamma
)
eager.register(
Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Poisson
)( # noqa: F821
eager.register(Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Poisson)( # noqa: F821
eager_gamma_poisson
)
if hasattr(dist, "DirichletMultinomial"):
eager.register(
Binary, ops.SubOp, JointDirichletMultinomial, DirichletMultinomial
)( # noqa: F821
eager.register(Binary, ops.SubOp, JointDirichletMultinomial, DirichletMultinomial)( # noqa: F821
eager_dirichlet_posterior
)
eager.register(
Reduce, ops.AddOp, Multinomial[Tensor, Funsor, Funsor], frozenset
)( # noqa: F821
eager.register(Reduce, ops.AddOp, Multinomial[Tensor, Funsor, Funsor], frozenset)( # noqa: F821
eager_plate_multinomial
)

Expand Down
12 changes: 6 additions & 6 deletions funsor/minipyro.py
Original file line number Diff line number Diff line change
Expand Up @@ -125,9 +125,9 @@ def __enter__(self):
# trace illustrates why we need postprocess_message in addition to process_message:
# We only want to record a value after all other effects have been applied
def postprocess_message(self, msg):
assert (
msg["type"] != "sample" or msg["name"] not in self.trace
), "sample sites must have unique names"
assert msg["type"] != "sample" or msg["name"] not in self.trace, (
"sample sites must have unique names"
)
self.trace[msg["name"]] = msg.copy()

def get_trace(self, *args, **kwargs):
Expand Down Expand Up @@ -238,9 +238,9 @@ def process_message(self, msg):

def postprocess_message(self, msg):
if msg["type"] == "sample":
assert (
msg["name"] not in self.log_factors
), "all sites must have unique names"
assert msg["name"] not in self.log_factors, (
"all sites must have unique names"
)
log_prob = msg["fn"].log_prob(msg["value"])
self.log_factors[msg["name"]] = log_prob
self.plates.update(f.name for f in msg["cond_indep_stack"].values())
Expand Down
6 changes: 3 additions & 3 deletions funsor/pyro/distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,9 +91,9 @@ def sample(self, sample_shape=torch.Size()):

def rsample(self, sample_shape=torch.Size()):
delta = self._sample_delta(sample_shape)
assert (
not delta.log_density.requires_grad
), "distribution is not fully reparametrized"
assert not delta.log_density.requires_grad, (
"distribution is not fully reparametrized"
)
ndims = len(sample_shape) + len(self.batch_shape) + len(self.event_shape)
value = funsor_to_tensor(delta.terms[0][1][0], ndims=ndims)
return value
Expand Down
6 changes: 3 additions & 3 deletions funsor/terms.py
Original file line number Diff line number Diff line change
Expand Up @@ -191,9 +191,9 @@ def __init__(cls, name, bases, dct):
def __getitem__(cls, arg_types):
if not isinstance(arg_types, tuple):
arg_types = (arg_types,)
assert len(arg_types) == len(
cls._ast_fields
), "Must provide exactly one type per subexpression"
assert len(arg_types) == len(cls._ast_fields), (
"Must provide exactly one type per subexpression"
)
return super().__getitem__(arg_types)

def __call__(cls, *args, **kwargs):
Expand Down
6 changes: 3 additions & 3 deletions funsor/testing.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,9 +115,9 @@ def assert_close(actual, expected, atol=1e-6, rtol=1e-6):
and isinstance(actual.terms[0], Tensor)
and is_array(actual.terms[0].data)
):
assert isinstance(expected, Contraction) and is_array(
expected.terms[0].data
), msg
assert isinstance(expected, Contraction) and is_array(expected.terms[0].data), (
msg
)
elif isinstance(actual, Contraction) and isinstance(actual.terms[0], Delta):
assert isinstance(expected, Contraction) and isinstance(
expected.terms[0], Delta
Expand Down
44 changes: 13 additions & 31 deletions funsor/torch/distributions.py
Original file line number Diff line number Diff line change
Expand Up @@ -336,9 +336,7 @@ def composetransform_to_funsor(tfm, output=None, dim_to_name=None, real_inputs=N
@to_funsor.register(torch.distributions.Bernoulli)
def bernoulli_to_funsor(pyro_dist, output=None, dim_to_name=None):
new_pyro_dist = _PyroWrapper_BernoulliLogits(logits=pyro_dist.logits)
return backenddist_to_funsor(
BernoulliLogits, new_pyro_dist, output, dim_to_name
) # noqa: F821
return backenddist_to_funsor(BernoulliLogits, new_pyro_dist, output, dim_to_name) # noqa: F821


@to_funsor.register(dist.Delta) # Delta **distribution**
Expand All @@ -358,25 +356,17 @@ def deltadist_to_funsor(pyro_dist, output=None, dim_to_name=None):
]


eager.register(Beta, Funsor, Funsor, Funsor)(eager_beta) # noqa: F821)
eager.register(Beta, Funsor, Funsor, Funsor)(eager_beta) # noqa: F821
eager.register(Binomial, Funsor, Funsor, Funsor)(eager_binomial) # noqa: F821
eager.register(Multinomial, Tensor, Tensor, Tensor)(eager_multinomial) # noqa: F821)
eager.register(Categorical, Funsor, Tensor)(eager_categorical_funsor) # noqa: F821)
eager.register(Categorical, Tensor, Variable)(eager_categorical_tensor) # noqa: F821)
eager.register(Categorical, Constant[Tuple, Tensor], Variable)(
eager_categorical_tensor
) # noqa: F821)
eager.register(Multinomial, Tensor, Tensor, Tensor)(eager_multinomial) # noqa: F821
eager.register(Categorical, Funsor, Tensor)(eager_categorical_funsor) # noqa: F821
eager.register(Categorical, Tensor, Variable)(eager_categorical_tensor) # noqa: F821
eager.register(Categorical, Constant[Tuple, Tensor], Variable)(eager_categorical_tensor) # noqa: F821
eager.register(Delta, Tensor, Tensor, Tensor)(eager_delta_tensor) # noqa: F821
eager.register(Delta, Funsor, Funsor, Variable)(
eager_delta_funsor_variable
) # noqa: F821
eager.register(Delta, Variable, Funsor, Variable)(
eager_delta_funsor_variable
) # noqa: F821
eager.register(Delta, Funsor, Funsor, Variable)(eager_delta_funsor_variable) # noqa: F821
eager.register(Delta, Variable, Funsor, Variable)(eager_delta_funsor_variable) # noqa: F821
eager.register(Delta, Variable, Funsor, Funsor)(eager_delta_funsor_funsor) # noqa: F821
eager.register(Delta, Variable, Variable, Variable)(
eager_delta_variable_variable
) # noqa: F821
eager.register(Delta, Variable, Variable, Variable)(eager_delta_variable_variable) # noqa: F821
eager.register(Normal, Funsor, Tensor, Funsor)(eager_normal) # noqa: F821
eager.register(MultivariateNormal, Funsor, Tensor, Funsor)(eager_mvn) # noqa: F821
eager.register(
Expand All @@ -394,23 +384,15 @@ def deltadist_to_funsor(pyro_dist, output=None, dim_to_name=None):
)( # noqa: F821
eager_dirichlet_multinomial
)
eager.register(
Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Gamma
)( # noqa: F821
eager.register(Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Gamma)( # noqa: F821
eager_gamma_gamma
)
eager.register(
Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Poisson
)( # noqa: F821
eager.register(Contraction, ops.LogaddexpOp, ops.AddOp, frozenset, Gamma, Poisson)( # noqa: F821
eager_gamma_poisson
)
eager.register(
Binary, ops.SubOp, JointDirichletMultinomial, DirichletMultinomial
)( # noqa: F821
eager.register(Binary, ops.SubOp, JointDirichletMultinomial, DirichletMultinomial)( # noqa: F821
eager_dirichlet_posterior
)
eager.register(
Reduce, ops.AddOp, Multinomial[Tensor, Funsor, Funsor], frozenset
)( # noqa: F821
eager.register(Reduce, ops.AddOp, Multinomial[Tensor, Funsor, Funsor], frozenset)( # noqa: F821
eager_plate_multinomial
)
6 changes: 3 additions & 3 deletions funsor/typing.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,9 +279,9 @@ def __getitem__(cls, arg_types):
assert not get_args(cls), "cannot subscript a subscripted type {}".format(
cls
)
assert not any(
isvariadic(arg_type) for arg_type in arg_types
), "nested variadic types not supported"
assert not any(isvariadic(arg_type) for arg_type in arg_types), (
"nested variadic types not supported"
)
new_dct = cls.__dict__.copy()
new_dct.update({"__args__": arg_types})
# type(cls) to handle GenericTypeMeta subclasses
Expand Down
5 changes: 3 additions & 2 deletions funsor/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,8 +257,9 @@ def set_backend(backend):
_JAX_COMPILED_FUNCTION_TYPE = type(jax.jit(lambda: 0))
else:
raise ValueError(
"backend should be either 'numpy', 'torch', or 'jax'"
", got {}".format(backend)
"backend should be either 'numpy', 'torch', or 'jax', got {}".format(
backend
)
)


Expand Down
Loading
Loading