開始使用免費開始

使用 Accelerator 的 Local SGD

你已經實作梯度累積(gradient accumulation)與梯度檢查點(gradient checkpointing),以精簡語言翻譯模型的記憶體使用。不過訓練仍有些慢,因此你決定在訓練迴圈中加入 local SGD,提升裝置之間的通訊效率。請用 local SGD 建立訓練迴圈!

modeltrain_dataloaderaccelerator 都已預先定義,且已匯入 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.____()
編輯並執行程式碼