Gradient checkpointing з Trainer
Ви хочете використати gradient checkpointing, щоб зменшити споживання памʼяті вашою моделлю. Ви вже бачили, як писати явний тренувальний цикл з Accelerator, а тепер хочете скористатися спрощеним інтерфейсом без циклів тренування через Trainer. Виконання вправи займе певний час через виклик trainer.train().
Налаштуйте аргументи для Trainer, щоб увімкнути gradient checkpointing.
Ця вправа є частиною курсу
Ефективне тренування моделей ШІ з PyTorch
Інструкції до вправи
- Використайте чотири кроки акумулювання градієнта в
TrainingArguments. - Увімкніть gradient checkpointing у
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()