(is_trace, use_mycb)
| 45 | @dist.launcher(n_gpus=2, device_type="gpu") |
| 46 | def worker(): |
| 47 | def runner(is_trace, use_mycb): |
| 48 | np.random.seed(dist.get_rank() + 123) |
| 49 | megengine.random.seed(dist.get_rank() + 123) |
| 50 | |
| 51 | model = ConvNet() |
| 52 | model.train() |
| 53 | |
| 54 | if dist.is_distributed(): |
| 55 | dist.bcast_list_(model.tensors()) |
| 56 | |
| 57 | side_effect_cnt = 0 |
| 58 | |
| 59 | def mycb(_, grad): |
| 60 | nonlocal side_effect_cnt |
| 61 | side_effect_cnt += 1 |
| 62 | return F.clip(grad, -2e-2, 2e-2) |
| 63 | |
| 64 | cblist = ( |
| 65 | [mycb, dist.make_allreduce_cb("mean")] |
| 66 | if use_mycb |
| 67 | else [dist.make_allreduce_cb("mean")] |
| 68 | ) |
| 69 | gm = autodiff.GradManager().attach(model.parameters(), callbacks=cblist) |
| 70 | optimizer = AdamW(model.parameters(), lr=0.01) |
| 71 | |
| 72 | image = np.random.randn(3, 8, 3, 32, 32) |
| 73 | label = np.random.randint(0, 10, (3, 8,)) |
| 74 | |
| 75 | def func(model, optimizer, timage, tlabel): |
| 76 | with gm: |
| 77 | score = model(timage) |
| 78 | loss = F.nn.cross_entropy(score, tlabel) |
| 79 | gm.backward(loss) |
| 80 | optimizer.step().clear_grad() |
| 81 | return loss |
| 82 | |
| 83 | if is_trace: |
| 84 | func = xla_trace(func, without_host=True, capture_as_const=True) |
| 85 | |
| 86 | losses, bn_states, opt_states, weights, ses = [], [], [], [], [] |
| 87 | for i in range(6): |
| 88 | timage = megengine.Tensor(image[i % 3]) |
| 89 | tlabel = megengine.Tensor(label[i % 3]) |
| 90 | loss = func(model, optimizer, timage, tlabel) |
| 91 | |
| 92 | losses.append(loss.item()) |
| 93 | bn_states.append(model.bn1.running_mean.numpy().reshape(-1)) |
| 94 | opt_states.append( |
| 95 | list(optimizer._state.values())[3]["exp_avg"].numpy().reshape(-1) |
| 96 | ) |
| 97 | weights.append(model.conv2.weight.numpy().reshape(-1)) |
| 98 | ses.append(side_effect_cnt) |
| 99 | |
| 100 | if i == 4: |
| 101 | for pg in optimizer.param_groups: |
| 102 | pg["lr"] = 0.006 |
| 103 | |
| 104 | return ( |
no test coverage detected