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
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()