勾配チェックポインティングとローカルSGD

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

Dennis Lee

Data Engineer, Amazon

学習効率の向上

 

 

メモリ効率、通信効率、計算効率を表すアイコン

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

勾配チェックポインティングによるメモリ効率の向上

 

 

メモリ効率、通信効率、計算効率を表すアイコン

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

ローカルSGDによる通信効率の改善

 

 

メモリ効率、通信効率、計算効率を表すアイコン

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算
    • まずA、Bを計算し、次にCを計算

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算
    • まずA、Bを計算し、次にCを計算
    • 順伝播の残りでA、Bは不要
  • A、Bを保存するか削除するか?

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算
    • まずA、Bを計算し、次にCを計算
    • 順伝播の残りでA、Bは不要
  • A、Bを保存するか削除するか?
    • チェックポインティングなし:A、Bを保存

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算
    • まずA、Bを計算し、次にCを計算
    • 順伝播の残りでA、Bは不要
  • A、Bを保存するか削除するか?
    • チェックポインティングなし:A、Bを保存
    • 勾配チェックポインティング:A、Bを削除

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算
    • まずA、Bを計算し、次にCを計算
    • 順伝播の残りでA、Bは不要
  • A、Bを保存するか削除するか?
    • チェックポインティングなし:A、Bを保存
    • 勾配チェックポインティング:A、Bを削除
    • 逆伝播時にA、Bを再計算

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

勾配チェックポインティングとは?

  • 勾配チェックポインティング:保存する活性化を選択してメモリを削減
  • 例:A + B = C を計算
    • まずA、Bを計算し、次にCを計算
    • 順伝播の残りでA、Bは不要
  • A、Bを保存するか削除するか?
    • チェックポインティングなし:A、Bを保存
    • 勾配チェックポインティング:A、Bを削除
    • 逆伝播時にA、Bを再計算
    • Bの再計算コストが高い場合は保存

ノードとエッジによる勾配チェックポインティングを示すグラフ

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

TrainerとAccelerator

AcceleratorとTrainerの使いやすさとカスタマイズ性を比較したグラフ

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

TrainerとAccelerator

AcceleratorとTrainerの使いやすさとカスタマイズ性を比較したグラフ

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

TrainerによるGradient Checkpointing

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







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

TrainerによるGradient Checkpointing

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

TrainerからAcceleratorへ

AcceleratorとTrainerの使いやすさとカスタマイズ性を比較したグラフ

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

AcceleratorによるGradient Checkpointing

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

AcceleratorによるGradient Checkpointing

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

ローカルSGDによる通信効率の向上

 

 

メモリ効率、通信効率、計算効率を表すアイコン

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

ローカルSGDとは?

一定ステップごとに勾配を同期するローカルSGDの動作を示す図

  • 各デバイスが並列に勾配を計算
PyTorch による効率的な AI モデルトレーニング

ローカルSGDとは?

一定ステップごとに勾配を同期するローカルSGDの動作を示す図

  • 各デバイスが並列に勾配を計算
  • 勾配の同期:ドライバノードが各デバイスのモデルパラメータを更新
  • ローカルSGD:勾配同期の頻度を削減
PyTorch による効率的な AI モデルトレーニング

Acceleratorによるローカル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()

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

AcceleratorによるローカルSGD

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

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

AcceleratorによるローカルSGD

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

練習しましょう!

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

Preparing Video For Download...