(BNModule, is_training, use_trace, use_symbolic)
| 13 | |
| 14 | |
| 15 | def run_frozen_bn(BNModule, is_training, use_trace, use_symbolic): |
| 16 | nchannel = 3 |
| 17 | m = BNModule(nchannel, freeze=True) |
| 18 | if is_training: |
| 19 | m.train() |
| 20 | else: |
| 21 | m.eval() |
| 22 | var = 4.0 |
| 23 | bias = 1.0 |
| 24 | shape = (1, nchannel, 1, 1) |
| 25 | m.running_var[...] = var * F.ones(shape) |
| 26 | m.running_mean[...] = bias * F.ones(shape) |
| 27 | |
| 28 | saved_var = m.running_var.numpy() |
| 29 | saved_mean = m.running_mean.numpy() |
| 30 | saved_wt = m.weight.numpy() |
| 31 | saved_bias = m.bias.numpy() |
| 32 | |
| 33 | gm = ad.GradManager().attach(m.parameters()) |
| 34 | optim = optimizer.SGD(m.parameters(), lr=1.0) |
| 35 | optim.clear_grad() |
| 36 | |
| 37 | data = np.random.random((6, nchannel, 2, 2)).astype("float32") |
| 38 | |
| 39 | def train_fn(d): |
| 40 | for _ in range(3): |
| 41 | with gm: |
| 42 | loss = m(d).mean() |
| 43 | gm.backward(loss) |
| 44 | optim.step() |
| 45 | return loss |
| 46 | |
| 47 | if use_trace: |
| 48 | train_fn = trace(train_fn, symbolic=use_symbolic) |
| 49 | |
| 50 | for _ in range(3): |
| 51 | loss = train_fn(megengine.tensor(data)) |
| 52 | if not is_training: |
| 53 | np.testing.assert_equal(m.running_var.numpy(), saved_var) |
| 54 | np.testing.assert_equal(m.running_mean.numpy(), saved_mean) |
| 55 | np.testing.assert_almost_equal( |
| 56 | loss.numpy(), ((data - bias) / np.sqrt(var)).mean(), 5 |
| 57 | ) |
| 58 | np.testing.assert_equal(m.weight.numpy(), saved_wt) |
| 59 | np.testing.assert_equal(m.bias.numpy(), saved_bias) |
| 60 | |
| 61 | |
| 62 | @pytest.mark.parametrize("is_training", [False, True]) |
no test coverage detected