Accelerator を使った勾配チェックポインティング
引き続き、使用するデバイス上で言語翻訳モデルをトレーニングできるよう、メモリ使用量の最適化を進めましょう。勾配累積(グラジェント・アキュムレーション)により、より大きなバッチサイズでの効果的なトレーニングが可能になりました。次のステップとして、モデルのメモリ使用量をさらに削減するために、勾配チェックポインティング(グラジェント・チェックポインティング)を追加しましょう。
model、train_dataloader、accelerator はあらかじめ定義されています。
この演習はコースの一部です
PyTorch による効率的な AI モデルトレーニング
演習の手順
modelの勾配チェックポインティングを有効にします。modelの勾配累積を有効にするためのAcceleratorコンテキストマネージャーを設定します。
実践的なインタラクティブ演習
このサンプルコードを完成させて、この演習に挑戦してみましょう。
# Enable gradient checkpointing on the model
____.____()
for batch in train_dataloader:
with accelerator.accumulate(model):
inputs, targets = batch["input_ids"], batch["labels"]
# Get the outputs from a forward pass of the model
____ = ____(____, labels=targets)
loss = outputs.loss
accelerator.backward(loss)
optimizer.step()
lr_scheduler.step()
optimizer.zero_grad()
print(f"Loss = {loss}")