Gradient checkpointing với Trainer
Bạn muốn dùng gradient checkpointing để giảm mức sử dụng bộ nhớ của mô hình. Bạn đã thấy cách viết vòng lặp huấn luyện tường minh với Accelerator, và giờ bạn muốn dùng giao diện đơn giản hơn, không cần vòng lặp huấn luyện, với Trainer. Bài tập sẽ mất một chút thời gian để chạy khi gọi trainer.train().
Thiết lập các đối số cho Trainer để sử dụng gradient checkpointing.
Bài tập này là một phần của khóa học
Huấn luyện Mô hình AI Hiệu quả với PyTorch
Hướng dẫn bài tập
- Dùng bốn bước tích lũy gradient trong
TrainingArguments. - Bật gradient checkpointing trong
TrainingArguments. - Truyền các đối số huấn luyện vào
Trainer.
Bài tập tương tác thực hành trực tiếp
Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.
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()