Addestramento a basso uso di memoria con Adafactor

Efficient AI Model Training with PyTorch

Dennis Lee

Data Engineer, Amazon

Ottimizzatori per un training efficiente

 

 

Icone che rappresentano AdamW, Adafactor e Adam a 8 bit.

Efficient AI Model Training with PyTorch

Compromessi degli ottimizzatori

Diagramma che mostra i compromessi tra numero di parametri e precisione per AdamW, Adafactor e Adam a 8 bit.

Efficient AI Model Training with PyTorch

Come funziona Adafactor?

 

Diagramma che mostra i passaggi di Adafactor.

Efficient AI Model Training with PyTorch

Come funziona Adafactor?

 

Diagramma che mostra i passaggi di Adafactor.

Efficient AI Model Training with PyTorch

Come funziona Adafactor?

 

Diagramma che mostra i passaggi di Adafactor.

 

  • EMA: media mobile esponenziale
  • Secondo momento: EMA dei gradienti al quadrato
Efficient AI Model Training with PyTorch

Come funziona Adafactor?

 

Diagramma che mostra i passaggi di Adafactor.

 

  • EMA: media mobile esponenziale
  • Secondo momento: EMA dei gradienti al quadrato
Efficient AI Model Training with PyTorch

Come Adafactor risparmia memoria?

 

Diagramma che mostra la matrice del secondo momento, somma per colonna e per riga.

  • Risparmia memoria evitando di salvare la matrice del secondo momento
Efficient AI Model Training with PyTorch

Come Adafactor risparmia memoria?

 

Diagramma che mostra la matrice del secondo momento, somma per colonna e per riga.

  • Risparmia memoria evitando di salvare la matrice del secondo momento
  • Invece salva le somme per colonna e per riga della matrice
Efficient AI Model Training with PyTorch

Come Adafactor risparmia memoria?

 

Diagramma che mostra la matrice del secondo momento, somma per colonna e per riga.

  • Risparmia memoria evitando di salvare la matrice del secondo momento
  • Invece salva le somme per colonna e per riga della matrice
Efficient AI Model Training with PyTorch

Come Adafactor risparmia memoria?

 

Diagramma che mostra la matrice del secondo momento, somma per colonna e per riga.

  • Risparmia memoria evitando di salvare la matrice del secondo momento
  • Invece salva le somme per colonna e per riga della matrice
  • Stima la matrice completa moltiplicando somma per colonna e per riga
Efficient AI Model Training with PyTorch

Implementazione con Trainer e Accelerator

Grafico che confronta facilità d’uso vs. possibilità di personalizzazione per Accelerator e Trainer.

Efficient AI Model Training with PyTorch

Implementa Adafactor con 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}
Efficient AI Model Training with PyTorch

Implementa Adafactor con Accelerator

# Assumes PyTorch 2.5 or higher
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
Efficient AI Model Training with PyTorch

Ispeziona lo stato dell’optimizer

  • Accedi allo state dell’optimizer
optimizer_state = optimizer.state.values()
  • Oppure accedi all’optimizer tramite 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, ...])}, ...])
Efficient AI Model Training with PyTorch

Calcola l’uso di memoria di Adafactor

total_size_megabytes, total_num_elements = \
    compute_optimizer_size(trainer.optimizer.state.values())
print(f"Number of Adafactor parameters: {total_num_elements:,}")
print(f"Adafactor size: {total_size_megabytes:.0f} MB")
Number of Adafactor parameters: 178,712
Adafactor size: 1 MB
  • Confronto con AdamW: Adafactor usa molta meno memoria!
Number of AdamW parameters: 131,566,188
AdamW size: 502 MB
Efficient AI Model Training with PyTorch

Passiamo alla pratica !

Efficient AI Model Training with PyTorch

Preparing Video For Download...