开始使用免费开始使用

结合 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()
编辑并运行代码