Loading src/borch/__init__.py +2 −2 Original line number Diff line number Diff line Loading @@ -63,8 +63,8 @@ Examples: from borch.version import __version__ from borch import infer, metrics, posterior from borch.random_variable import RandomVariable, RVPair from borch import infer, metrics, posterior, distributions 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/module.py +9 −7 Original line number Diff line number Diff line Loading @@ -292,6 +292,7 @@ class Observed(_Module): ) class Module(_Module): """Acts as a ``torch.nn.Module`` but handles ``borch.RandomVariable`` s correctly. Loading Loading @@ -433,7 +434,9 @@ class Module(_Module): msg = "Invalid arguments: only None or kwargs allowed" raise ValueError(msg) for key in self.observed: getattr(self.posterior, key)() # redraw this sample rv = getattr(self.posterior, key, None) if rv is not None: rv() # redraw this sample self.observed.clear() not_tensor_or_none = [ Loading @@ -458,18 +461,17 @@ class Module(_Module): return out except AttributeError: posterior = self.__dict__.get("_modules", {}).get("posterior", None) prior = self.__dict__.get("_modules", {}).get("prior", None) if posterior is not None: param = getattr(self.posterior, name, None) if param is None and isinstance( getattr(self.prior, name, None), borch.RandomVariable getattr(prior, name, None), borch.RandomVariable ): self.posterior.set_random_variable( name, getattr(self.prior, name), None ) param = getattr(self.posterior, name) self.posterior.set_random_variable(name, getattr(prior, name), None) param = getattr(posterior, name) if isinstance(param, borch.RandomVariable): observed = self.observed.get(name) if self.observed.get(name) is not None: if observed is not None: param.tensor = observed self._used_rvs.add(name) # we convert it to just a tensor here to avoid potential bugs Loading src/borch/nn/__init__.py +1 −1 Original line number Diff line number Diff line Loading @@ -63,5 +63,5 @@ Notes: from borch.nn.torch_proxies import * from borch.nn import utils from borch.nn.borchify import borchify_module, borchify_network from borch.nn.borchify import borchify_module, borchify_network, borchify_namespace from borch.module import Module Loading
src/borch/__init__.py +2 −2 Original line number Diff line number Diff line Loading @@ -63,8 +63,8 @@ Examples: from borch.version import __version__ from borch import infer, metrics, posterior from borch.random_variable import RandomVariable, RVPair from borch import infer, metrics, posterior, distributions 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/module.py +9 −7 Original line number Diff line number Diff line Loading @@ -292,6 +292,7 @@ class Observed(_Module): ) class Module(_Module): """Acts as a ``torch.nn.Module`` but handles ``borch.RandomVariable`` s correctly. Loading Loading @@ -433,7 +434,9 @@ class Module(_Module): msg = "Invalid arguments: only None or kwargs allowed" raise ValueError(msg) for key in self.observed: getattr(self.posterior, key)() # redraw this sample rv = getattr(self.posterior, key, None) if rv is not None: rv() # redraw this sample self.observed.clear() not_tensor_or_none = [ Loading @@ -458,18 +461,17 @@ class Module(_Module): return out except AttributeError: posterior = self.__dict__.get("_modules", {}).get("posterior", None) prior = self.__dict__.get("_modules", {}).get("prior", None) if posterior is not None: param = getattr(self.posterior, name, None) if param is None and isinstance( getattr(self.prior, name, None), borch.RandomVariable getattr(prior, name, None), borch.RandomVariable ): self.posterior.set_random_variable( name, getattr(self.prior, name), None ) param = getattr(self.posterior, name) self.posterior.set_random_variable(name, getattr(prior, name), None) param = getattr(posterior, name) if isinstance(param, borch.RandomVariable): observed = self.observed.get(name) if self.observed.get(name) is not None: if observed is not None: param.tensor = observed self._used_rvs.add(name) # we convert it to just a tensor here to avoid potential bugs Loading
src/borch/nn/__init__.py +1 −1 Original line number Diff line number Diff line Loading @@ -63,5 +63,5 @@ Notes: from borch.nn.torch_proxies import * from borch.nn import utils from borch.nn.borchify import borchify_module, borchify_network from borch.nn.borchify import borchify_module, borchify_network, borchify_namespace from borch.module import Module