開始使用免費開始

使用 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()
編輯並執行程式碼