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

add transformer tests

parent a4fd88c8
Loading
Loading
Loading
Loading
+25 −0
Original line number Diff line number Diff line
@@ -261,3 +261,28 @@ def test_sequential():
    assert isinstance(net(torch.ones(1, 3)), torch.Tensor)
    loss = vi_loss(**pq_to_infer(net))
    assert loss > 0

def test_torch_save_whole_module():
    lin = borch.nn.Linear(2, 1)
    buffer = BytesIO()
    torch.save(lin, buffer)
    buffer.seek(0)
    new_lin = torch.load(buffer)
    assert isinstance(new_lin, borch.Module)
    out = new_lin(torch.ones(2))
    borch.sample(new_lin)
    assert out != new_lin(torch.ones(2))


def test_transformer_samples():
    transformer_model = borch.nn.Transformer(nhead=16, num_encoder_layers=12)
    src = torch.rand((10, 32, 512))
    tgt = torch.rand((20, 32, 512))
    out = transformer_model(src, tgt).sum()
    borch.sample(transformer_model)
    assert out != transformer_model(src, tgt).sum()

def test_transformer_submodules_are_borch_modules():
    transformer_model = borch.nn.Transformer(nhead=16, num_encoder_layers=12)
    for name, mod in transformer_model.named_child_modules():
        import ipdb; ipdb.set_trace()