Loading tests/borch/nn/utils/test_borchify.py +36 −1 Original line number Diff line number Diff line from unittest import TestCase, mock import numpy.testing as npt import torch from borch import nn from borch import nn, sample from borch.nn import borchify from borch.posterior import ScaledNormal from borch.module import random_variables Loading Loading @@ -187,3 +188,37 @@ def test_borchify_a_borch_module(): new = borchify.borchify_network(net, posterior_creator=lambda: ScaledNormal(1e-6)) assert isinstance(new(torch.randn(1, 2)), torch.Tensor) assert isinstance(new.posterior, ScaledNormal) class MultiplyOrDevide(torch.nn.Module): def __init__(self, multiply=False): super().__init__() self.multiply = multiply self.parameter = torch.nn.Parameter(torch.tensor([1, 2, 3.0])) def forward(self, x): if self.multiply: return self.parameter * x return self.parameter / x def test_custom_module_multiply(): mul = borchify.borchify_network( MultiplyOrDevide(True), posterior_creator=lambda: ScaledNormal(1e-2, loc_at_prior_mean=False), ) assert mul.multiply x = torch.tensor([2.0]) out = mul(x) sample(mul) assert all(out != mul(x)) npt.assert_allclose(out.detach(), [2, 4, 6], rtol=0.1) def test_custom_module_devide(): devide = borchify.borchify_network( MultiplyOrDevide(False), posterior_creator=lambda: ScaledNormal(1e-2, loc_at_prior_mean=False), ) assert not devide.multiply npt.assert_allclose(devide(torch.tensor(2.0)).detach(), [0.5, 1, 1.5], rtol=0.1) Loading
tests/borch/nn/utils/test_borchify.py +36 −1 Original line number Diff line number Diff line from unittest import TestCase, mock import numpy.testing as npt import torch from borch import nn from borch import nn, sample from borch.nn import borchify from borch.posterior import ScaledNormal from borch.module import random_variables Loading Loading @@ -187,3 +188,37 @@ def test_borchify_a_borch_module(): new = borchify.borchify_network(net, posterior_creator=lambda: ScaledNormal(1e-6)) assert isinstance(new(torch.randn(1, 2)), torch.Tensor) assert isinstance(new.posterior, ScaledNormal) class MultiplyOrDevide(torch.nn.Module): def __init__(self, multiply=False): super().__init__() self.multiply = multiply self.parameter = torch.nn.Parameter(torch.tensor([1, 2, 3.0])) def forward(self, x): if self.multiply: return self.parameter * x return self.parameter / x def test_custom_module_multiply(): mul = borchify.borchify_network( MultiplyOrDevide(True), posterior_creator=lambda: ScaledNormal(1e-2, loc_at_prior_mean=False), ) assert mul.multiply x = torch.tensor([2.0]) out = mul(x) sample(mul) assert all(out != mul(x)) npt.assert_allclose(out.detach(), [2, 4, 6], rtol=0.1) def test_custom_module_devide(): devide = borchify.borchify_network( MultiplyOrDevide(False), posterior_creator=lambda: ScaledNormal(1e-2, loc_at_prior_mean=False), ) assert not devide.multiply npt.assert_allclose(devide(torch.tensor(2.0)).detach(), [0.5, 1, 1.5], rtol=0.1)