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

Function run_dtr_resnet1202

imperative/python/test/integration/test_dtr.py:105–140  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

103
104
105def run_dtr_resnet1202():
106 batch_size = 6
107 resnet1202 = ResNet(BasicBlock, [200, 200, 200])
108 opt = optim.SGD(resnet1202.parameters(), lr=0.05, momentum=0.9, weight_decay=1e-4)
109 gm = GradManager().attach(resnet1202.parameters())
110
111 def train_func(data, label, *, net, gm):
112 net.train()
113 with gm:
114 pred = net(data)
115 loss = F.loss.cross_entropy(pred, label)
116 gm.backward(loss)
117 return pred, loss
118
119 _, free_mem = mge.device.get_mem_status_bytes()
120 tensor_mem = free_mem - (2 ** 30)
121 if tensor_mem > 0:
122 x = np.ones((1, int(tensor_mem / 4)), dtype=np.float32)
123 else:
124 x = np.ones((1,), dtype=np.float32)
125 t = mge.tensor(x)
126
127 mge.dtr.enable()
128 mge.dtr.enable_sqrt_sampling = True
129
130 data = np.random.randn(batch_size, 3, 32, 32).astype("float32")
131 label = np.random.randint(0, 10, size=(batch_size,)).astype("int32")
132 for _ in range(2):
133 opt.clear_grad()
134 _, loss = train_func(mge.tensor(data), mge.tensor(label), net=resnet1202, gm=gm)
135 opt.step()
136 loss.item()
137
138 t.numpy()
139 mge.dtr.disable()
140 mge._exit(0)
141
142
143@pytest.mark.require_ngpu(1)

Callers

nothing calls this directly

Calls 15

GradManagerClass · 0.90
parametersMethod · 0.80
onesMethod · 0.80
tensorMethod · 0.80
ResNetClass · 0.70
train_funcFunction · 0.70
SGDMethod · 0.45
attachMethod · 0.45
get_mem_status_bytesMethod · 0.45
enableMethod · 0.45
astypeMethod · 0.45
clear_gradMethod · 0.45

Tested by

no test coverage detected