开始使用免费开始使用

使用 Accelerator 进行梯度检查点

您正在继续优化内存使用,以便在本地设备上训练机器翻译模型。梯度累积已经帮助您用更大的批量有效训练。在此基础上加入梯度检查点,以进一步降低模型的内存占用。

modeltrain_dataloaderaccelerator 已预先定义。

本练习是课程的一部分

使用 PyTorch 高效训练 AI 模型

查看课程

练习说明

  • model 上启用梯度检查点。
  • 设置 Accelerator 的上下文管理器,在 model 上启用梯度累积。

交互式实操练习

通过完成这段示例代码来试试这个练习。

# Enable gradient checkpointing on the model
____.____()

for batch in train_dataloader:
    with accelerator.accumulate(model):
        inputs, targets = batch["input_ids"], batch["labels"]
        # Get the outputs from a forward pass of the model
        ____ = ____(____, labels=targets)
        loss = outputs.loss
        accelerator.backward(loss)
        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()
        print(f"Loss = {loss}")
编辑并运行代码