Checkpointing de gradientes y SGD local

Entrenamiento eficiente de modelos de IA con PyTorch

Dennis Lee

Data Engineer, Amazon

Mejorar la eficiencia del entrenamiento

 

 

Iconos que representan eficiencia de memoria, comunicación y cómputo.

Entrenamiento eficiente de modelos de IA con PyTorch

El checkpointing de gradientes mejora la eficiencia de memoria

 

 

Iconos que representan eficiencia de memoria, comunicación y cómputo.

Entrenamiento eficiente de modelos de IA con PyTorch

La SGD local aborda la eficiencia de comunicación

 

 

Iconos que representan eficiencia de memoria, comunicación y cómputo.

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C
    • Primero calcula A y B, luego C

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C
    • Primero calcula A y B, luego C
    • A y B no se necesitan para el resto del forward
  • ¿Guardamos o eliminamos A y B?

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C
    • Primero calcula A y B, luego C
    • A y B no se necesitan para el resto del forward
  • ¿Guardamos o eliminamos A y B?
    • Sin checkpointing: guarda A y B

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C
    • Primero calcula A y B, luego C
    • A y B no se necesitan para el resto del forward
  • ¿Guardamos o eliminamos A y B?
    • Sin checkpointing: guarda A y B
    • Con checkpointing: elimina A y B

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C
    • Primero calcula A y B, luego C
    • A y B no se necesitan para el resto del forward
  • ¿Guardamos o eliminamos A y B?
    • Sin checkpointing: guarda A y B
    • Con checkpointing: elimina A y B
    • Recalcula A y B en el backward

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es el checkpointing de gradientes?

  • Checkpointing de gradientes: reduce memoria eligiendo qué activaciones guardar
  • Ejemplo: calcula A + B = C
    • Primero calcula A y B, luego C
    • A y B no se necesitan para el resto del forward
  • ¿Guardamos o eliminamos A y B?
    • Sin checkpointing: guarda A y B
    • Con checkpointing: elimina A y B
    • Recalcula A y B en el backward
    • Si B es caro de recalcular, guárdalo

Gráfico que ilustra el checkpointing de gradientes con nodos y aristas

Entrenamiento eficiente de modelos de IA con PyTorch

Trainer y Accelerator

Gráfico que compara facilidad de uso vs. capacidad de personalización para Accelerator y Trainer.

Entrenamiento eficiente de modelos de IA con PyTorch

Trainer y Accelerator

Gráfico que compara facilidad de uso vs. capacidad de personalización para Accelerator y Trainer.

Entrenamiento eficiente de modelos de IA con PyTorch

Checkpointing de gradientes con Trainer

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







Entrenamiento eficiente de modelos de IA con PyTorch

Checkpointing de gradientes 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}
Entrenamiento eficiente de modelos de IA con PyTorch

De Trainer a Accelerator

Gráfico que compara facilidad de uso vs. capacidad de personalización para Accelerator y Trainer.

Entrenamiento eficiente de modelos de IA con PyTorch

Checkpointing de gradientes 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()
Entrenamiento eficiente de modelos de IA con PyTorch

Checkpointing de gradientes 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()
Entrenamiento eficiente de modelos de IA con PyTorch

La SGD local mejora la eficiencia de comunicación

 

 

Iconos que representan eficiencia de memoria, comunicación y cómputo.

Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es la SGD local?

Diagrama que muestra cómo funciona la SGD local sincronizando gradientes tras cierto número de pasos.

  • Cada dispositivo calcula gradientes en paralelo
Entrenamiento eficiente de modelos de IA con PyTorch

¿Qué es la SGD local?

Diagrama que muestra cómo funciona la SGD local sincronizando gradientes tras cierto número de pasos.

  • Cada dispositivo calcula gradientes en paralelo
  • Sincronización de gradientes: el nodo driver actualiza los parámetros en cada dispositivo
  • SGD local: reduce la frecuencia de sincronización
Entrenamiento eficiente de modelos de IA con PyTorch

SGD local 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()

Entrenamiento eficiente de modelos de IA con PyTorch

SGD local 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()

Entrenamiento eficiente de modelos de IA con PyTorch

SGD local 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()
Entrenamiento eficiente de modelos de IA con PyTorch

¡Vamos a practicar!

Entrenamiento eficiente de modelos de IA con PyTorch

Preparing Video For Download...