使用 Accelerator 的 Local SGD
你已經實作梯度累積(gradient accumulation)與梯度檢查點(gradient checkpointing),以精簡語言翻譯模型的記憶體使用。不過訓練仍有些慢,因此你決定在訓練迴圈中加入 local SGD,提升裝置之間的通訊效率。請用 local SGD 建立訓練迴圈!
model、train_dataloader 與 accelerator 都已預先定義,且已匯入 LocalSGD。
本練習屬於課程
使用 PyTorch 高效訓練 AI 模型
練習說明
- 將
local_sgd_steps設為每 8 個步驟同步一次梯度。 - 讓 local SGD 內容管理器前進一步(step)。
動手互動練習
試著完成這個範例程式碼,體驗一下這個練習。
# Set up a context manager to synchronize gradients every eight steps
with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=____, enabled=True) as local_sgd:
for batch in train_dataloader:
with accelerator.accumulate(model):
inputs, targets = batch["input_ids"], batch["labels"]
outputs = model(inputs, labels=targets)
loss = outputs.loss
accelerator.backward(loss)
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad()
# Step the local SGD context manager
local_sgd.____()