Zacznij terazZacznij za darmo

Gradient checkpointing z Trainerem

Chcesz użyć gradient checkpointingu, aby zmniejszyć zużycie pamięci przez model. Wiesz już, jak napisać pętlę treningową z użyciem Accelerator – teraz skorzystasz z uproszczonego interfejsu Trainer, który nie wymaga ręcznego kodowania pętli. Pamiętaj, że wywołanie trainer.train() może chwilę potrwać.

Skonfiguruj argumenty Trainer, aby korzystał z gradient checkpointingu.

To ćwiczenie jest częścią kursu

Efektywne trenowanie modeli AI z PyTorch

Zobacz kurs

Instrukcje do ćwiczenia

  • Ustaw cztery kroki akumulacji gradientu w TrainingArguments.
  • Włącz gradient checkpointing w TrainingArguments.
  • Przekaż argumenty treningowe do Trainer.

Interaktywne ćwiczenie praktyczne

Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.

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()
Edytuj i uruchom kod