开始使用免费开始使用

使用 Accelerator 的本地 SGD

您已经实现了梯度累积和梯度检查点,以优化翻译模型的内存使用。训练仍然有些慢,因此您决定在训练循环中加入本地 SGD,以提升设备间的通信效率。请使用本地 SGD 构建训练循环!

modeltrain_dataloaderaccelerator 已预定义,且已导入 LocalSGD

本练习是课程的一部分

使用 PyTorch 高效训练 AI 模型

查看课程

练习说明

  • local_sgd_steps 设为每 8 步同步梯度。
  • 调用本地 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.____()
编辑并运行代码