ПочатиПочніть безкоштовно

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()
Редагувати та запускати код