Entrenamiento eficiente en memoria con Adafactor

Entrenamiento eficiente de modelos de IA con PyTorch

Dennis Lee

Data Engineer, Amazon

Optimizadores para entrenar con eficiencia

 

 

Iconos que representan AdamW, Adafactor y Adam de 8 bits.

Entrenamiento eficiente de modelos de IA con PyTorch

Compensaciones de optimizadores

Diagrama que muestra las compensaciones entre número de parámetros y precisión para AdamW, Adafactor y Adam de 8 bits.

Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo funciona Adafactor?

 

Diagrama que muestra los pasos de Adafactor.

Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo funciona Adafactor?

 

Diagrama que muestra los pasos de Adafactor.

Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo funciona Adafactor?

 

Diagrama que muestra los pasos de Adafactor.

 

  • EMA: media móvil exponencial
  • Segundo momento: EMA de los gradientes al cuadrado
Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo funciona Adafactor?

 

Diagrama que muestra los pasos de Adafactor.

 

  • EMA: media móvil exponencial
  • Segundo momento: EMA de los gradientes al cuadrado
Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo ahorra memoria Adafactor?

 

Diagrama que muestra la matriz de segundo momento, la suma por columna y la suma por fila.

  • Ahorra memoria sin guardar la matriz de segundo momento
Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo ahorra memoria Adafactor?

 

Diagrama que muestra la matriz de segundo momento, la suma por columna y la suma por fila.

  • Ahorra memoria sin guardar la matriz de segundo momento
  • En su lugar, guarda la suma por columna y por fila
Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo ahorra memoria Adafactor?

 

Diagrama que muestra la matriz de segundo momento, la suma por columna y la suma por fila.

  • Ahorra memoria sin guardar la matriz de segundo momento
  • En su lugar, guarda la suma por columna y por fila
Entrenamiento eficiente de modelos de IA con PyTorch

¿Cómo ahorra memoria Adafactor?

 

Diagrama que muestra la matriz de segundo momento, la suma por columna y la suma por fila.

  • Ahorra memoria sin guardar la matriz de segundo momento
  • En su lugar, guarda la suma por columna y por fila
  • Estima la matriz completa multiplicando ambas sumas
Entrenamiento eficiente de modelos de IA con PyTorch

Implementación con 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

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

Implementar 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
Entrenamiento eficiente de modelos de IA con PyTorch

Inspeccionar el estado del optimizador

  • Accede al optimizer mediante su state
optimizer_state = optimizer.state.values()
  • O accede al optimizer a través de 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, ...])}, ...])
Entrenamiento eficiente de modelos de IA con PyTorch

Calcular el uso de memoria de 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
  • Comparado con AdamW: ¡Adafactor usa mucha menos memoria!
Number of AdamW parameters: 131,566,188
AdamW size: 502 MB
Entrenamiento eficiente de modelos de IA con PyTorch

¡Vamos a practicar!

Entrenamiento eficiente de modelos de IA con PyTorch

Preparing Video For Download...