MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / test_bn_no_track_stat2

Function test_bn_no_track_stat2

imperative/python/test/integration/test_bn.py:97–119  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

95
96
97def 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
122def test_bn_no_track_stat3():

Callers

nothing calls this directly

Calls 12

BatchNorm2dClass · 0.90
parametersMethod · 0.80
assert_equalMethod · 0.80
numpyMethod · 0.45
attachMethod · 0.45
GradManagerMethod · 0.45
SGDMethod · 0.45
clear_gradMethod · 0.45
astypeMethod · 0.45
sumMethod · 0.45
backwardMethod · 0.45
stepMethod · 0.45

Tested by

no test coverage detected