Gradient-Checkpointing und lokales SGD

Effizientes KI-Modelltraining mit PyTorch

Dennis Lee

Data Engineer, Amazon

Trainingseffizienz steigern

 

 

Symbole für Speicher-, Kommunikations- und Recheneffizienz.

Effizientes KI-Modelltraining mit PyTorch

Gradient-Checkpointing verbessert Speichereffizienz

 

 

Symbole für Speicher-, Kommunikations- und Recheneffizienz.

Effizientes KI-Modelltraining mit PyTorch

Lokales SGD verbessert Kommunikationseffizienz

 

 

Symbole für Speicher-, Kommunikations- und Recheneffizienz.

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C
    • Erst A, B berechnen, dann C

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C
    • Erst A, B berechnen, dann C
    • A, B werden im restlichen Forward-Pass nicht mehr benötigt
  • Sollen A und B gespeichert oder verworfen werden?

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C
    • Erst A, B berechnen, dann C
    • A, B werden im restlichen Forward-Pass nicht mehr benötigt
  • Sollen A und B gespeichert oder verworfen werden?
    • Ohne Gradient-Checkpointing: A, B speichern

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C
    • Erst A, B berechnen, dann C
    • A, B werden im restlichen Forward-Pass nicht mehr benötigt
  • Sollen A und B gespeichert oder verworfen werden?
    • Ohne Gradient-Checkpointing: A, B speichern
    • Mit Gradient-Checkpointing: A, B verwerfen

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C
    • Erst A, B berechnen, dann C
    • A, B werden im restlichen Forward-Pass nicht mehr benötigt
  • Sollen A und B gespeichert oder verworfen werden?
    • Ohne Gradient-Checkpointing: A, B speichern
    • Mit Gradient-Checkpointing: A, B verwerfen
    • A, B im Backward-Pass neu berechnen

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Was ist Gradient-Checkpointing?

  • Gradient-Checkpointing: Speicher sparen, indem du auswählst, welche Aktivierungen gespeichert werden
  • Beispiel: berechne A + B = C
    • Erst A, B berechnen, dann C
    • A, B werden im restlichen Forward-Pass nicht mehr benötigt
  • Sollen A und B gespeichert oder verworfen werden?
    • Ohne Gradient-Checkpointing: A, B speichern
    • Mit Gradient-Checkpointing: A, B verwerfen
    • A, B im Backward-Pass neu berechnen
    • Wenn B teuer ist, B speichern

Grafik zum Gradient-Checkpointing mit Knoten und Kanten

Effizientes KI-Modelltraining mit PyTorch

Trainer und Accelerator

Diagramm: Bedienkomfort vs. Anpassbarkeit für Accelerator und Trainer.

Effizientes KI-Modelltraining mit PyTorch

Trainer und Accelerator

Diagramm: Bedienkomfort vs. Anpassbarkeit für Accelerator und Trainer.

Effizientes KI-Modelltraining mit PyTorch

Gradient-Checkpointing mit Trainer

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







Effizientes KI-Modelltraining mit PyTorch

Gradient-Checkpointing mit 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}
Effizientes KI-Modelltraining mit PyTorch

Vom Trainer zum Accelerator

Diagramm: Bedienkomfort vs. Anpassbarkeit für Accelerator und Trainer.

Effizientes KI-Modelltraining mit PyTorch

Gradient-Checkpointing mit 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()
Effizientes KI-Modelltraining mit PyTorch

Gradient-Checkpointing mit 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()
Effizientes KI-Modelltraining mit PyTorch

Lokales SGD verbessert Kommunikationseffizienz

 

 

Symbole für Speicher-, Kommunikations- und Recheneffizienz.

Effizientes KI-Modelltraining mit PyTorch

Was ist lokales SGD?

Diagramm, das zeigt, wie lokales SGD durch Synchronisieren nach einigen Schritten funktioniert.

  • Jedes Gerät berechnet Gradienten parallel
Effizientes KI-Modelltraining mit PyTorch

Was ist lokales SGD?

Diagramm, das zeigt, wie lokales SGD durch Synchronisieren nach einigen Schritten funktioniert.

  • Jedes Gerät berechnet Gradienten parallel
  • Gradienten-Synchronisierung: Der Driver-Knoten aktualisiert Modellparameter auf jedem Gerät
  • Lokales SGD: Häufigkeit der Gradienten-Synchronisierung verringern
Effizientes KI-Modelltraining mit PyTorch

Lokales SGD mit 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()

Effizientes KI-Modelltraining mit PyTorch

Lokales SGD mit 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()

Effizientes KI-Modelltraining mit PyTorch

Lokales SGD mit 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()
Effizientes KI-Modelltraining mit PyTorch

Lass uns üben!

Effizientes KI-Modelltraining mit PyTorch

Preparing Video For Download...