Adafactorによるメモリ効率の高い学習

PyTorch による効率的な AI モデルトレーニング

Dennis Lee

Data Engineer, Amazon

学習効率のためのオプティマイザー

 

 

AdamW、Adafactor、8-bit Adamを表すアイコン。

PyTorch による効率的な AI モデルトレーニング

オプティマイザーのトレードオフ

AdamW、Adafactor、8-bit Adamのパラメーター数と精度のトレードオフを示す図。

PyTorch による効率的な AI モデルトレーニング

Adafactorの仕組み

 

Adafactorのステップを示す図。

PyTorch による効率的な AI モデルトレーニング

Adafactorの仕組み

 

Adafactorのステップを示す図。

PyTorch による効率的な AI モデルトレーニング

Adafactorの仕組み

 

Adafactorのステップを示す図。

 

  • EMA:指数移動平均
  • 二次モーメント:勾配の二乗のEMA
PyTorch による効率的な AI モデルトレーニング

Adafactorの仕組み

 

Adafactorのステップを示す図。

 

  • EMA:指数移動平均
  • 二次モーメント:勾配の二乗のEMA
PyTorch による効率的な AI モデルトレーニング

Adafactorのメモリ節約の仕組み

 

二次モーメント行列、列和、行和を示す図。

  • 二次モーメント行列を保存しないことでメモリを節約
PyTorch による効率的な AI モデルトレーニング

Adafactorのメモリ節約の仕組み

 

二次モーメント行列、列和、行和を示す図。

  • 二次モーメント行列を保存しないことでメモリを節約
  • 代わりに行列の列和と行和を保存
PyTorch による効率的な AI モデルトレーニング

Adafactorのメモリ節約の仕組み

 

二次モーメント行列、列和、行和を示す図。

  • 二次モーメント行列を保存しないことでメモリを節約
  • 代わりに行列の列和と行和を保存
PyTorch による効率的な AI モデルトレーニング

Adafactorのメモリ節約の仕組み

 

二次モーメント行列、列和、行和を示す図。

  • 二次モーメント行列を保存しないことでメモリを節約
  • 代わりに行列の列和と行和を保存
  • 列和と行和の積で完全な行列を近似
PyTorch による効率的な AI モデルトレーニング

TrainerとAcceleratorの実装

AcceleratorとTrainerの使いやすさとカスタマイズ性を比較したチャート。

PyTorch による効率的な AI モデルトレーニング

TrainerでAdafactorを実装する

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}
PyTorch による効率的な AI モデルトレーニング

AcceleratorでAdafactorを実装する

# 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
PyTorch による効率的な AI モデルトレーニング

オプティマイザーの状態を確認する

  • optimizerstate を通じてアクセス
optimizer_state = optimizer.state.values()
  • または trainer 経由で optimizer にアクセス
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, ...])}, ...])
PyTorch による効率的な AI モデルトレーニング

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
  • AdamWと比較:Adafactorは大幅にメモリ使用量が少ない!
Number of AdamW parameters: 131,566,188
AdamW size: 502 MB
PyTorch による効率的な AI モデルトレーニング

練習しましょう!

PyTorch による効率的な AI モデルトレーニング

Preparing Video For Download...