Kom igångKom igång gratis

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

Visa kurs

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