Commit 230a02e4 authored by Johan Gudmundsson's avatar Johan Gudmundsson
Browse files

borchify test of custom module

parent bfe11781
Loading
Loading
Loading
Loading
+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
@@ -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)