Loading src/borch/__init__.py +1 −1 Original line number Diff line number Diff line Loading @@ -64,7 +64,7 @@ Examples: from borch.version import __version__ from borch import infer, metrics, posterior from borch.random_variable import RandomVariable, RVPair from borch.random_variable import RandomVariable, RVPair, validate_args from borch.graph import Graph, as_tensor from borch.module import ( Module, Loading src/borch/distributions/distributions.py +9 −5 Original line number Diff line number Diff line Loading @@ -24,9 +24,9 @@ class Delta(distributions.Distribution): support = constraints.real arg_constraints = {"value": constraints.real} def __init__(self, value, event_shape=torch.Size()): def __init__(self, value, event_shape=torch.Size(), validate_args=None): self.value = as_tensor(value) super().__init__(self.value.size(), event_shape) super().__init__(self.value.size(), event_shape, validate_args=validate_args) @property def mean(self): Loading Loading @@ -64,9 +64,13 @@ class PointMass(distributions.TransformedDistribution): arg_constraints = {"value": constraints.real} def __init__(self, value, support, event_shape=torch.Size()): base_dist = Delta(as_tensor(value), event_shape=event_shape) super(PointMass, self).__init__(base_dist, transform_to(support)) def __init__(self, value, support, event_shape=torch.Size(), validate_args=None): base_dist = Delta( as_tensor(value), event_shape=event_shape, validate_args=validate_args ) super(PointMass, self).__init__( base_dist, transform_to(support), validate_args=validate_args ) @property def value(self): Loading src/borch/distributions/rv_distributions.py +29 −23 Original line number Diff line number Diff line Loading @@ -13,7 +13,7 @@ from borch.utils.func_tools import disable_doctests class _LocScaleArgs(RandomVariable): def __init__(self, loc, scale, validate_args=False, posterior=None): def __init__(self, loc, scale, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("loc", loc) self.register_param_or_buffer("scale", scale) Loading @@ -25,7 +25,7 @@ class _LocScaleArgs(RandomVariable): class _ScaleArg(RandomVariable): def __init__(self, scale, validate_args=False, posterior=None): def __init__(self, scale, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("scale", scale) Loading @@ -46,7 +46,7 @@ class HalfNormal(_ScaleArg): class _ProbsLogitsArgs(RandomVariable): def __init__(self, probs=None, logits=None, validate_args=False, posterior=None): def __init__(self, probs=None, logits=None, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) self.register_param_or_buffer("probs", probs) Loading Loading @@ -111,7 +111,7 @@ class StudentT(RandomVariable): distribution_cls = _dist.StudentT __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, df, loc, scale, validate_args=False, posterior=None): def __init__(self, df, loc, scale, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("df", df) self.register_param_or_buffer("loc", loc) Loading @@ -130,7 +130,7 @@ class Pareto(RandomVariable): distribution_cls = _dist.Pareto __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, scale, alpha, validate_args=False, posterior=None): def __init__(self, scale, alpha, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("scale", scale) self.register_param_or_buffer("alpha", alpha) Loading @@ -151,7 +151,7 @@ def _check_only_one(**kwargs): class _Binomial(RandomVariable): def __init__( self, total_count, probs=None, logits=None, validate_args=False, posterior=None self, total_count, probs=None, logits=None, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) Loading Loading @@ -183,7 +183,7 @@ class Multinomial(RandomVariable): __doc__ = disable_doctests(distribution_cls.__doc__) def __init__( self, total_count, probs=None, logits=None, validate_args=False, posterior=None self, total_count, probs=None, logits=None, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) Loading @@ -204,7 +204,7 @@ class Gamma(RandomVariable): distribution_cls = _dist.Gamma __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, concentration, rate, validate_args=False, posterior=None): def __init__(self, concentration, rate, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration", concentration) self.register_param_or_buffer("rate", rate) Loading @@ -221,7 +221,7 @@ class LKJCholesky(RandomVariable): distribution_cls = _dist.LKJCholesky __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, dimension, concentration, validate_args=False, posterior=None): def __init__(self, dimension, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration", concentration) self.register_param_or_buffer("dimension", dimension) Loading @@ -238,7 +238,7 @@ class Dirichlet(RandomVariable): distribution_cls = _dist.Dirichlet __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, concentration, validate_args=False, posterior=None): def __init__(self, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration", concentration) Loading @@ -249,7 +249,7 @@ class Dirichlet(RandomVariable): class _RateArg(RandomVariable): def __init__(self, rate, validate_args=False, posterior=None): def __init__(self, rate, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("rate", rate) Loading @@ -273,7 +273,7 @@ class Chi2(RandomVariable): distribution_cls = _dist.Chi2 __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, df, validate_args=False, posterior=None): def __init__(self, df, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("df", df) Loading @@ -287,18 +287,20 @@ class FisherSnedecor(RandomVariable): distribution_cls = _dist.FisherSnedecor __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, df1, df2, validate_args=False, posterior=None): def __init__(self, df1, df2, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("df1", df1) self.register_param_or_buffer("df2", df2) def _distribution(self): return self.distribution_cls(as_tensor(self.df1), as_tensor(self.df2)) return self.distribution_cls( as_tensor(self.df1), as_tensor(self.df2), validate_args=self.validate_args ) class _Concentraton1Concentration0Args(RandomVariable): def __init__( self, concentration1, concentration0, validate_args=False, posterior=None self, concentration1, concentration0, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration0", concentration0) Loading Loading @@ -326,7 +328,7 @@ class Uniform(RandomVariable): distribution_cls = _dist.Uniform __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, low, high, validate_args=False, posterior=None): def __init__(self, low, high, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("low", low) self.register_param_or_buffer("high", high) Loading @@ -347,7 +349,9 @@ class PointMass(RandomVariable): self._support = support def _distribution(self): return self.distribution_cls(as_tensor(self.value), self._support) return self.distribution_cls( as_tensor(self.value), self._support, validate_args=self.validate_args ) class Delta(RandomVariable): Loading @@ -359,7 +363,9 @@ class Delta(RandomVariable): self.register_param_or_buffer("value", value) def _distribution(self): return self.distribution_cls(as_tensor(self.value)) return self.distribution_cls( as_tensor(self.value), validate_args=self.validate_args ) class TransformedDistribution(RandomVariable): Loading @@ -367,7 +373,7 @@ class TransformedDistribution(RandomVariable): __doc__ = disable_doctests(distribution_cls.__doc__) def __init__( self, base_distribution, transforms, validate_args=False, posterior=None self, base_distribution, transforms, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) # NB use add_module instead of setattr, as we don't want this added to Loading @@ -393,7 +399,7 @@ class MultivariateNormal(RandomVariable): covariance_matrix=None, precision_matrix=None, scale_tril=None, validate_args=False, validate_args=None, posterior=None, ): _check_only_one( Loading @@ -419,7 +425,7 @@ class MultivariateNormal(RandomVariable): class _TemperatureProbsLogitsArgs(RandomVariable): def __init__( self, temperature, probs=None, logits=None, validate_args=False, posterior=None self, temperature, probs=None, logits=None, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) Loading Loading @@ -450,7 +456,7 @@ class VonMises(RandomVariable): distribution_cls = _dist.VonMises __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, loc, concentration, validate_args=False, posterior=None): def __init__(self, loc, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("loc", loc) self.register_param_or_buffer("concentration", concentration) Loading @@ -467,7 +473,7 @@ class Weibull(RandomVariable): distribution_cls = _dist.Weibull __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, scale, concentration, validate_args=False, posterior=None): def __init__(self, scale, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("scale", scale) self.register_param_or_buffer("concentration", concentration) Loading src/borch/random_variable.py +43 −2 Original line number Diff line number Diff line Loading @@ -2,11 +2,36 @@ The base class for the RandomVariable primitive. """ import contextlib import contextvars from typing import Optional, Any, Union import torch import borch.graph as graph _VALIDATE_ARGS = contextvars.ContextVar("VALIDATE_ARGS", default=None) @contextlib.contextmanager def validate_args(value): """ Context manager that sets the `validate_args` for all random variable distributions. """ _correct_validate_args_value(value) token = _VALIDATE_ARGS.set(value) try: yield finally: _VALIDATE_ARGS.reset(token) def _correct_validate_args_value(value): if value is not None and not isinstance(value, bool): raise ValueError( f"`validate_args` must be one of True, False or None not {value}" ) class RandomVariable(graph.Graph): """Base class for a ``RandomVariable`` primitive used to model stochastic nodes. Loading Loading @@ -53,9 +78,25 @@ class RandomVariable(graph.Graph): tensor(1., requires_grad=True)] """ def __init__(self, validate_args=False, posterior=None): def __init__( self, validate_args=None, posterior=None ): # pylint: disable=redefined-outer-name _correct_validate_args_value(validate_args) super().__init__(posterior=posterior) self.validate_args = validate_args self._validate_args = validate_args @property def validate_args(self): """if the args should be validated when creating the distribution""" validate = self._validate_args if self._validate_args is not None else __debug__ ctx_value = _VALIDATE_ARGS.get() return ctx_value if ctx_value is not None else validate @validate_args.setter def validate_args(self, value): """Set the value for validate_args""" _correct_validate_args_value(value) self._validate_args = value def __repr__(self): params = {**self._buffers, **self._modules, **self._parameters} Loading tests/borch/distributions/test_rv_distributions.py +7 −0 Original line number Diff line number Diff line Loading @@ -114,3 +114,10 @@ def test_can_set_posterior(dist): kwargs["posterior"] = borch.posterior.ScaledNormal() new_dist = type(dist)(**kwargs) assert isinstance(new_dist.posterior, borch.posterior.ScaledNormal) def test_validate_args_ctxt(dist): with borch.validate_args(True): assert dist.validate_args with borch.validate_args(False): assert not dist.validate_args Loading
src/borch/__init__.py +1 −1 Original line number Diff line number Diff line Loading @@ -64,7 +64,7 @@ Examples: from borch.version import __version__ from borch import infer, metrics, posterior from borch.random_variable import RandomVariable, RVPair from borch.random_variable import RandomVariable, RVPair, validate_args from borch.graph import Graph, as_tensor from borch.module import ( Module, Loading
src/borch/distributions/distributions.py +9 −5 Original line number Diff line number Diff line Loading @@ -24,9 +24,9 @@ class Delta(distributions.Distribution): support = constraints.real arg_constraints = {"value": constraints.real} def __init__(self, value, event_shape=torch.Size()): def __init__(self, value, event_shape=torch.Size(), validate_args=None): self.value = as_tensor(value) super().__init__(self.value.size(), event_shape) super().__init__(self.value.size(), event_shape, validate_args=validate_args) @property def mean(self): Loading Loading @@ -64,9 +64,13 @@ class PointMass(distributions.TransformedDistribution): arg_constraints = {"value": constraints.real} def __init__(self, value, support, event_shape=torch.Size()): base_dist = Delta(as_tensor(value), event_shape=event_shape) super(PointMass, self).__init__(base_dist, transform_to(support)) def __init__(self, value, support, event_shape=torch.Size(), validate_args=None): base_dist = Delta( as_tensor(value), event_shape=event_shape, validate_args=validate_args ) super(PointMass, self).__init__( base_dist, transform_to(support), validate_args=validate_args ) @property def value(self): Loading
src/borch/distributions/rv_distributions.py +29 −23 Original line number Diff line number Diff line Loading @@ -13,7 +13,7 @@ from borch.utils.func_tools import disable_doctests class _LocScaleArgs(RandomVariable): def __init__(self, loc, scale, validate_args=False, posterior=None): def __init__(self, loc, scale, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("loc", loc) self.register_param_or_buffer("scale", scale) Loading @@ -25,7 +25,7 @@ class _LocScaleArgs(RandomVariable): class _ScaleArg(RandomVariable): def __init__(self, scale, validate_args=False, posterior=None): def __init__(self, scale, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("scale", scale) Loading @@ -46,7 +46,7 @@ class HalfNormal(_ScaleArg): class _ProbsLogitsArgs(RandomVariable): def __init__(self, probs=None, logits=None, validate_args=False, posterior=None): def __init__(self, probs=None, logits=None, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) self.register_param_or_buffer("probs", probs) Loading Loading @@ -111,7 +111,7 @@ class StudentT(RandomVariable): distribution_cls = _dist.StudentT __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, df, loc, scale, validate_args=False, posterior=None): def __init__(self, df, loc, scale, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("df", df) self.register_param_or_buffer("loc", loc) Loading @@ -130,7 +130,7 @@ class Pareto(RandomVariable): distribution_cls = _dist.Pareto __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, scale, alpha, validate_args=False, posterior=None): def __init__(self, scale, alpha, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("scale", scale) self.register_param_or_buffer("alpha", alpha) Loading @@ -151,7 +151,7 @@ def _check_only_one(**kwargs): class _Binomial(RandomVariable): def __init__( self, total_count, probs=None, logits=None, validate_args=False, posterior=None self, total_count, probs=None, logits=None, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) Loading Loading @@ -183,7 +183,7 @@ class Multinomial(RandomVariable): __doc__ = disable_doctests(distribution_cls.__doc__) def __init__( self, total_count, probs=None, logits=None, validate_args=False, posterior=None self, total_count, probs=None, logits=None, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) Loading @@ -204,7 +204,7 @@ class Gamma(RandomVariable): distribution_cls = _dist.Gamma __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, concentration, rate, validate_args=False, posterior=None): def __init__(self, concentration, rate, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration", concentration) self.register_param_or_buffer("rate", rate) Loading @@ -221,7 +221,7 @@ class LKJCholesky(RandomVariable): distribution_cls = _dist.LKJCholesky __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, dimension, concentration, validate_args=False, posterior=None): def __init__(self, dimension, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration", concentration) self.register_param_or_buffer("dimension", dimension) Loading @@ -238,7 +238,7 @@ class Dirichlet(RandomVariable): distribution_cls = _dist.Dirichlet __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, concentration, validate_args=False, posterior=None): def __init__(self, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration", concentration) Loading @@ -249,7 +249,7 @@ class Dirichlet(RandomVariable): class _RateArg(RandomVariable): def __init__(self, rate, validate_args=False, posterior=None): def __init__(self, rate, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("rate", rate) Loading @@ -273,7 +273,7 @@ class Chi2(RandomVariable): distribution_cls = _dist.Chi2 __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, df, validate_args=False, posterior=None): def __init__(self, df, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("df", df) Loading @@ -287,18 +287,20 @@ class FisherSnedecor(RandomVariable): distribution_cls = _dist.FisherSnedecor __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, df1, df2, validate_args=False, posterior=None): def __init__(self, df1, df2, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("df1", df1) self.register_param_or_buffer("df2", df2) def _distribution(self): return self.distribution_cls(as_tensor(self.df1), as_tensor(self.df2)) return self.distribution_cls( as_tensor(self.df1), as_tensor(self.df2), validate_args=self.validate_args ) class _Concentraton1Concentration0Args(RandomVariable): def __init__( self, concentration1, concentration0, validate_args=False, posterior=None self, concentration1, concentration0, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("concentration0", concentration0) Loading Loading @@ -326,7 +328,7 @@ class Uniform(RandomVariable): distribution_cls = _dist.Uniform __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, low, high, validate_args=False, posterior=None): def __init__(self, low, high, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("low", low) self.register_param_or_buffer("high", high) Loading @@ -347,7 +349,9 @@ class PointMass(RandomVariable): self._support = support def _distribution(self): return self.distribution_cls(as_tensor(self.value), self._support) return self.distribution_cls( as_tensor(self.value), self._support, validate_args=self.validate_args ) class Delta(RandomVariable): Loading @@ -359,7 +363,9 @@ class Delta(RandomVariable): self.register_param_or_buffer("value", value) def _distribution(self): return self.distribution_cls(as_tensor(self.value)) return self.distribution_cls( as_tensor(self.value), validate_args=self.validate_args ) class TransformedDistribution(RandomVariable): Loading @@ -367,7 +373,7 @@ class TransformedDistribution(RandomVariable): __doc__ = disable_doctests(distribution_cls.__doc__) def __init__( self, base_distribution, transforms, validate_args=False, posterior=None self, base_distribution, transforms, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) # NB use add_module instead of setattr, as we don't want this added to Loading @@ -393,7 +399,7 @@ class MultivariateNormal(RandomVariable): covariance_matrix=None, precision_matrix=None, scale_tril=None, validate_args=False, validate_args=None, posterior=None, ): _check_only_one( Loading @@ -419,7 +425,7 @@ class MultivariateNormal(RandomVariable): class _TemperatureProbsLogitsArgs(RandomVariable): def __init__( self, temperature, probs=None, logits=None, validate_args=False, posterior=None self, temperature, probs=None, logits=None, validate_args=None, posterior=None ): super().__init__(validate_args=validate_args, posterior=posterior) _check_only_one(probs=probs, logits=logits) Loading Loading @@ -450,7 +456,7 @@ class VonMises(RandomVariable): distribution_cls = _dist.VonMises __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, loc, concentration, validate_args=False, posterior=None): def __init__(self, loc, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("loc", loc) self.register_param_or_buffer("concentration", concentration) Loading @@ -467,7 +473,7 @@ class Weibull(RandomVariable): distribution_cls = _dist.Weibull __doc__ = disable_doctests(distribution_cls.__doc__) def __init__(self, scale, concentration, validate_args=False, posterior=None): def __init__(self, scale, concentration, validate_args=None, posterior=None): super().__init__(validate_args=validate_args, posterior=posterior) self.register_param_or_buffer("scale", scale) self.register_param_or_buffer("concentration", concentration) Loading
src/borch/random_variable.py +43 −2 Original line number Diff line number Diff line Loading @@ -2,11 +2,36 @@ The base class for the RandomVariable primitive. """ import contextlib import contextvars from typing import Optional, Any, Union import torch import borch.graph as graph _VALIDATE_ARGS = contextvars.ContextVar("VALIDATE_ARGS", default=None) @contextlib.contextmanager def validate_args(value): """ Context manager that sets the `validate_args` for all random variable distributions. """ _correct_validate_args_value(value) token = _VALIDATE_ARGS.set(value) try: yield finally: _VALIDATE_ARGS.reset(token) def _correct_validate_args_value(value): if value is not None and not isinstance(value, bool): raise ValueError( f"`validate_args` must be one of True, False or None not {value}" ) class RandomVariable(graph.Graph): """Base class for a ``RandomVariable`` primitive used to model stochastic nodes. Loading Loading @@ -53,9 +78,25 @@ class RandomVariable(graph.Graph): tensor(1., requires_grad=True)] """ def __init__(self, validate_args=False, posterior=None): def __init__( self, validate_args=None, posterior=None ): # pylint: disable=redefined-outer-name _correct_validate_args_value(validate_args) super().__init__(posterior=posterior) self.validate_args = validate_args self._validate_args = validate_args @property def validate_args(self): """if the args should be validated when creating the distribution""" validate = self._validate_args if self._validate_args is not None else __debug__ ctx_value = _VALIDATE_ARGS.get() return ctx_value if ctx_value is not None else validate @validate_args.setter def validate_args(self, value): """Set the value for validate_args""" _correct_validate_args_value(value) self._validate_args = value def __repr__(self): params = {**self._buffers, **self._modules, **self._parameters} Loading
tests/borch/distributions/test_rv_distributions.py +7 −0 Original line number Diff line number Diff line Loading @@ -114,3 +114,10 @@ def test_can_set_posterior(dist): kwargs["posterior"] = borch.posterior.ScaledNormal() new_dist = type(dist)(**kwargs) assert isinstance(new_dist.posterior, borch.posterior.ScaledNormal) def test_validate_args_ctxt(dist): with borch.validate_args(True): assert dist.validate_args with borch.validate_args(False): assert not dist.validate_args