Gradientkontrollpunkter med Trainer
Du vill använda gradientkontrollpunkter för att minska modellens minnesanvändning. Du har sett hur man skriver den explicita träningsloopen med Accelerator, och nu vill du använda ett förenklat gränssnitt utan träningsloopar med Trainer. Övningen tar lite tid att köra när trainer.train() anropas.
Konfigurera argumenten för Trainer så att gradientkontrollpunkter används.
Den här övningen är en del av kursen
Effektiv AI-modellträning med PyTorch
Övningsinstruktioner
- Använd fyra gradientackumuleringssteg i
TrainingArguments. - Aktivera gradientkontrollpunkter i
TrainingArguments. - Skicka in träningsargumenten till
Trainer.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
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()