Gradientkontrollpunkter med Accelerator
Du fortsätter att optimera minnesanvändningen så att du kan träna din språköversättningsmodell på din enhet. Gradientackumulering har hjälpt dig att effektivt träna med större batchstorlekar. Bygg vidare på det arbetet och lägg till gradientkontrollpunkter för att minska modellens minnesbehov.
model, train_dataloader och accelerator är fördefinierade.
Den här övningen är en del av kursen
Effektiv AI-modellträning med PyTorch
Övningsinstruktioner
- Aktivera gradientkontrollpunkter på
model. - Sätt upp en
Accelerator-kontexthanterare för att aktivera gradientackumulering påmodel.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
# Enable gradient checkpointing on the model
____.____()
for batch in train_dataloader:
with accelerator.accumulate(model):
inputs, targets = batch["input_ids"], batch["labels"]
# Get the outputs from a forward pass of the model
____ = ____(____, labels=targets)
loss = outputs.loss
accelerator.backward(loss)
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad()
print(f"Loss = {loss}")