结合 Trainer 的梯度检查点
您希望使用梯度检查点来降低模型的内存占用。您已经了解了如何用 Accelerator 编写显式训练循环,现在想使用无需手写训练循环的简化接口 Trainer。由于会调用 trainer.train(),本练习运行需要一些时间。
请设置 Trainer 的参数以启用梯度检查点。
本练习是课程的一部分
使用 PyTorch 高效训练 AI 模型
练习说明
- 在
TrainingArguments中使用 4 个梯度累积步数。 - 在
TrainingArguments中启用梯度检查点。 - 将训练参数传递给
Trainer。
交互式实操练习
通过完成这段示例代码来试试这个练习。
training_args = TrainingArguments(output_dir="./results",
evaluation_strategy="epoch",
# Use four gradient accumulation steps
gradient_accumulation_steps=____,
# Enable gradient checkpointing
gradient_checkpointing=____)
trainer = Trainer(model=model,
# Pass in the training arguments
args=____,
train_dataset=dataset["train"],
eval_dataset=dataset["validation"],
compute_metrics=compute_metrics)
trainer.train()