Commit e8970248 authored by Johan Gudmundsson's avatar Johan Gudmundsson
Browse files

more tests for torch proxies

parents a838c722 812f8110
Loading
Loading
Loading
Loading
+2 −2
Original line number Diff line number Diff line
@@ -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,
+9 −5
Original line number Diff line number Diff line
@@ -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):
@@ -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):
+29 −23
Original line number Diff line number Diff line
@@ -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)
@@ -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)

@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)
@@ -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)

@@ -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)

@@ -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)

@@ -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)
@@ -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)
@@ -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):
@@ -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):
@@ -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
@@ -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(
@@ -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)
@@ -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)
@@ -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)
+9 −7
Original line number Diff line number Diff line
@@ -292,6 +292,7 @@ class Observed(_Module):
        )



class Module(_Module):
    """Acts as a ``torch.nn.Module`` but handles ``borch.RandomVariable`` s correctly.

@@ -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 = [
@@ -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
+1 −1
Original line number Diff line number Diff line
@@ -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