Test `.forward()` doesn't throw an error
(net, img_size)
| 28 | |
| 29 | @pytest.mark.parametrize('img_size', [224, 256, 512]) |
| 30 | def test_forward(net, img_size): |
| 31 | """Test `.forward()` doesn't throw an error""" |
| 32 | data = torch.zeros((1, 3, img_size, img_size)) |
| 33 | output = net(data) |
| 34 | assert not torch.isnan(output).any() |
| 35 | |
| 36 | |
| 37 | def test_dropout_training(net): |