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

validate_args context manager and add missing validate_args proppagation

parent a4fd88c8
Loading
Loading
Loading
Loading
+1 −1
Original line number Diff line number Diff line
@@ -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,
+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)
+43 −2
Original line number Diff line number Diff line
@@ -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.
@@ -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}
+7 −0
Original line number Diff line number Diff line
@@ -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