Kom igångKom igång gratis

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

Visa kurs

Ö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()
Redigera och kör kod