()
| 95 | |
| 96 | |
| 97 | def test_bn_no_track_stat2(): |
| 98 | nchannel = 3 |
| 99 | m = BatchNorm2d(nchannel) # Init with track_running_stat = True |
| 100 | m.track_running_stats = False |
| 101 | |
| 102 | # m.running_var and m.running_mean created during init time |
| 103 | saved_var = m.running_var.numpy() |
| 104 | assert saved_var is not None |
| 105 | saved_mean = m.running_mean.numpy() |
| 106 | assert saved_mean is not None |
| 107 | |
| 108 | gm = ad.GradManager().attach(m.parameters()) |
| 109 | optim = optimizer.SGD(m.parameters(), lr=1.0) |
| 110 | optim.clear_grad() |
| 111 | |
| 112 | data = tensor(np.random.random((6, nchannel, 2, 2)).astype("float32")) |
| 113 | with gm: |
| 114 | loss = m(data).sum() |
| 115 | gm.backward(loss) |
| 116 | optim.step() |
| 117 | |
| 118 | np.testing.assert_equal(m.running_var.numpy(), saved_var) |
| 119 | np.testing.assert_equal(m.running_mean.numpy(), saved_mean) |
| 120 | |
| 121 | |
| 122 | def test_bn_no_track_stat3(): |
nothing calls this directly
no test coverage detected