使用 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()