Gradient checkpointing e local SGD

Efficient AI Model Training with PyTorch

Dennis Lee

Data Engineer, Amazon

Migliorare l'efficienza dell'addestramento

 

 

Icone che rappresentano efficienza di memoria, comunicazione e calcolo.

Efficient AI Model Training with PyTorch

Il gradient checkpointing migliora l'efficienza di memoria

 

 

Icone che rappresentano efficienza di memoria, comunicazione e calcolo.

Efficient AI Model Training with PyTorch

La local SGD migliora l'efficienza di comunicazione

 

 

Icone che rappresentano efficienza di memoria, comunicazione e calcolo.

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C
    • Prima calcola A e B, poi calcola C

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C
    • Prima calcola A e B, poi calcola C
    • A e B non servono per il resto del forward pass
  • Salviamo o rimuoviamo A e B?

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C
    • Prima calcola A e B, poi calcola C
    • A e B non servono per il resto del forward pass
  • Salviamo o rimuoviamo A e B?
    • Senza gradient checkpointing: salva A e B

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C
    • Prima calcola A e B, poi calcola C
    • A e B non servono per il resto del forward pass
  • Salviamo o rimuoviamo A e B?
    • Senza gradient checkpointing: salva A e B
    • Con gradient checkpointing: rimuovi A e B

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C
    • Prima calcola A e B, poi calcola C
    • A e B non servono per il resto del forward pass
  • Salviamo o rimuoviamo A e B?
    • Senza gradient checkpointing: salva A e B
    • Con gradient checkpointing: rimuovi A e B
    • Ricalcola A e B nel backward pass

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Cos'è il gradient checkpointing?

  • Gradient checkpointing: riduce la memoria scegliendo quali attivazioni salvare
  • Esempio: calcola A + B = C
    • Prima calcola A e B, poi calcola C
    • A e B non servono per il resto del forward pass
  • Salviamo o rimuoviamo A e B?
    • Senza gradient checkpointing: salva A e B
    • Con gradient checkpointing: rimuovi A e B
    • Ricalcola A e B nel backward pass
    • Se B è costoso da ricalcolare, salvalo

Grafico che illustra il gradient checkpointing con nodi e archi

Efficient AI Model Training with PyTorch

Trainer e Accelerator

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

Efficient AI Model Training with PyTorch

Trainer e Accelerator

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

Efficient AI Model Training with PyTorch

Gradient checkpointing con Trainer

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







Efficient AI Model Training with PyTorch

Gradient checkpointing con Trainer

training_args = TrainingArguments(output_dir="./results",
                                  evaluation_strategy="epoch",
                                  gradient_accumulation_steps=4,
                                  gradient_checkpointing=True)

trainer = Trainer(model=model, args=training_args, train_dataset=dataset["train"], eval_dataset=dataset["validation"], compute_metrics=compute_metrics)
trainer.train()
{'epoch': 1.0, 'eval_loss': 0.73, 'eval_accuracy': 0.03, 'eval_f1': 0.05}
Efficient AI Model Training with PyTorch

Da Trainer ad Accelerator

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

Efficient AI Model Training with PyTorch

Gradient checkpointing con Accelerator

accelerator = Accelerator(gradient_accumulation_steps=2)


for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        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()
Efficient AI Model Training with PyTorch

Gradient checkpointing con Accelerator

accelerator = Accelerator(gradient_accumulation_steps=2)
model.gradient_checkpointing_enable()

for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        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()
Efficient AI Model Training with PyTorch

La local SGD migliora l'efficienza di comunicazione

 

 

Icone che rappresentano efficienza di memoria, comunicazione e calcolo.

Efficient AI Model Training with PyTorch

Cos'è la local SGD?

Diagramma che mostra come la local SGD sincronizza i gradienti dopo un certo numero di step.

  • Ogni dispositivo calcola i gradienti in parallelo
Efficient AI Model Training with PyTorch

Cos'è la local SGD?

Diagramma che mostra come la local SGD sincronizza i gradienti dopo un certo numero di step.

  • Ogni dispositivo calcola i gradienti in parallelo
  • Sincronizzazione dei gradienti: il nodo driver aggiorna i parametri del modello su ogni dispositivo
  • Local SGD: riduci la frequenza di sincronizzazione dei gradienti
Efficient AI Model Training with PyTorch

Local SGD con Accelerator





for index, batch in enumerate(dataloader):
    with accelerator.accumulate(model):
        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()

Efficient AI Model Training with PyTorch

Local SGD con Accelerator

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            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()

Efficient AI Model Training with PyTorch

Local SGD con Accelerator

from accelerate.local_sgd import LocalSGD

with LocalSGD(accelerator=accelerator, model=model, local_sgd_steps=8, 
              enabled=True) as local_sgd:
    for index, batch in enumerate(dataloader):
        with accelerator.accumulate(model):
            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()
            local_sgd.step()
Efficient AI Model Training with PyTorch

Ayo berlatih!

Efficient AI Model Training with PyTorch

Preparing Video For Download...