使用 Accelerator 进行梯度检查点
您正在继续优化内存使用,以便在本地设备上训练机器翻译模型。梯度累积已经帮助您用更大的批量有效训练。在此基础上加入梯度检查点,以进一步降低模型的内存占用。
model、train_dataloader 和 accelerator 已预先定义。
本练习是课程的一部分
使用 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}")