Geheugenzuinig trainen met Adafactor

Efficiënt AI-modellen trainen met PyTorch

Dennis Lee

Data Engineer, Amazon

Optimizers voor efficiënter trainen

 

 

Pictogrammen voor AdamW, Adafactor en 8-bit Adam.

Efficiënt AI-modellen trainen met PyTorch

Afwegingen bij optimizers

Diagram met de afweging tussen aantal parameters en precisie voor AdamW, Adafactor en 8-bit Adam.

Efficiënt AI-modellen trainen met PyTorch

Hoe werkt Adafactor?

 

Diagram met de stappen van Adafactor.

Efficiënt AI-modellen trainen met PyTorch

Hoe werkt Adafactor?

 

Diagram met de stappen van Adafactor.

Efficiënt AI-modellen trainen met PyTorch

Hoe werkt Adafactor?

 

Diagram met de stappen van Adafactor.

 

  • EMA: exponentieel voortschrijdend gemiddelde
  • Tweede moment: EMA van de gekwadrateerde gradiënten
Efficiënt AI-modellen trainen met PyTorch

Hoe werkt Adafactor?

 

Diagram met de stappen van Adafactor.

 

  • EMA: exponentieel voortschrijdend gemiddelde
  • Tweede moment: EMA van de gekwadrateerde gradiënten
Efficiënt AI-modellen trainen met PyTorch

Hoe bespaart Adafactor geheugen?

 

Diagram met de tweede-momentmatrix, kolomsom en rijsom.

  • Bespaar geheugen door de tweede-momentmatrix niet op te slaan
Efficiënt AI-modellen trainen met PyTorch

Hoe bespaart Adafactor geheugen?

 

Diagram met de tweede-momentmatrix, kolomsom en rijsom.

  • Bespaar geheugen door de tweede-momentmatrix niet op te slaan
  • Sla in plaats daarvan de kolomsom en rijsom van de matrix op
Efficiënt AI-modellen trainen met PyTorch

Hoe bespaart Adafactor geheugen?

 

Diagram met de tweede-momentmatrix, kolomsom en rijsom.

  • Bespaar geheugen door de tweede-momentmatrix niet op te slaan
  • Sla in plaats daarvan de kolomsom en rijsom van de matrix op
Efficiënt AI-modellen trainen met PyTorch

Hoe bespaart Adafactor geheugen?

 

Diagram met de tweede-momentmatrix, kolomsom en rijsom.

  • Bespaar geheugen door de tweede-momentmatrix niet op te slaan
  • Sla in plaats daarvan de kolomsom en rijsom van de matrix op
  • Schat de volledige matrix door kolomsom en rijsom te vermenigvuldigen
Efficiënt AI-modellen trainen met PyTorch

Implementatie met Trainer en Accelerator

Grafiek die gebruiksgemak vs. aanpasbaarheid vergelijkt voor Accelerator en Trainer.

Efficiënt AI-modellen trainen met PyTorch

Adafactor implementeren met Trainer

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",

optim="adafactor")
trainer = Trainer(model=model, args=training_args, train_dataset=train_dataset, eval_dataset=validation_dataset, compute_metrics=compute_metrics) trainer.train()
{'epoch': 1.0, 'eval_accuracy': 0.6, 'eval_f1': 0.5}
Efficiënt AI-modellen trainen met PyTorch

Adafactor implementeren met Accelerator

# Vereist PyTorch 2.5 of hoger
from torch.optim import Adafactor

optimizer = Adafactor(params=model.parameters(), lr=lr)
for batch in train_dataloader: inputs, targets = batch["input_ids"], batch["labels"] outputs = model(inputs, labels=targets) loss = outputs.loss accelerator.backward(loss) optimizer.step() lr_scheduler.step() optimizer.zero_grad() print(f"Loss = {loss}")
Loss = 0.71
Efficiënt AI-modellen trainen met PyTorch

De optimizerstatus inspecteren

  • Benader de optimizer via zijn state
optimizer_state = optimizer.state.values()
  • Of benader de optimizer via trainer
optimizer_state = trainer.optimizer.state.values()
print(optimizer_state)
dict_values([{'step': tensor(3.),
              'exp_avg_sq_row': tensor([1.0000e-30, 1.0000e-30, 1.0000e-30,  ...]), 
              'exp_avg_sq_col': tensor([4.6147e-11, 5.5115e-11, 1.6338e-10, ...])}, ...])
Efficiënt AI-modellen trainen met PyTorch

Geheugengebruik van Adafactor berekenen

total_size_megabytes, total_num_elements = \
    compute_optimizer_size(trainer.optimizer.state.values())
print(f"Aantal Adafactor-parameters: {total_num_elements:,}")
print(f"Adafactor-grootte: {total_size_megabytes:.0f} MB")
Aantal Adafactor-parameters: 178,712
Adafactor-grootte: 1 MB
  • Vergelijk met AdamW: Adafactor gebruikt veel minder geheugen!
Aantal AdamW-parameters: 131,566,188
AdamW-grootte: 502 MB
Efficiënt AI-modellen trainen met PyTorch

Laten we oefenen!

Efficiënt AI-modellen trainen met PyTorch

Preparing Video For Download...