def train_batch(X, y, device_params, devices, lr):
X_shards, y_shards = split_batch(X, y, devices)
ls = [loss(lenet(Xs, dev_W), ys).sum()
for Xs, ys, dev_W in zip(X_shards, y_shards, device_params)]
for l in ls:
l.backward()
with torch.no_grad():
for i in range(len(device_params[0])):
allreduce([device_params[c][i].grad for c in range(len(devices))])
for param in device_params:
d2l.sgd(param, lr, X.shape[0])
for p in param:
p.grad = None