()
| 103 | |
| 104 | |
| 105 | def 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) |
nothing calls this directly
no test coverage detected