AI・機械学習
上級

AdamW(アダムダブリュー)

Adam オプティマイザの改良版で、Weight Decay(重み減衰)を勾配更新から分離して正則化を適切に機能させる手法。2019 年に Loshchilov & Hutter が提案し、GPT-4・Llama 3・Mistral Large など 2024-2026 年の主要 LLM すべてが採用する事実上の標準オプティマイザ。

0 回閲覧
0 いいね

AdamW とは

AdamW は、2019 年に Ilya Loshchilov と Frank Hutter が論文「Decoupled Weight Decay Regularization」で提案したオプティマイザである。従来の Adam では Weight Decay が適応的学習率と結合し、正則化効果が不安定になる問題(L2 正則化と Weight Decay の非等価性)があった。AdamW はこの結合を解除(Decouple)し、Weight Decay を独立したステップとして適用することで、広い学習率範囲で安定した正則化を実現する。

Adam との違い

Adam と AdamW の核心的な違いは Weight Decay の適用方法にある:

  • Adam + L2 正則化: 勾配に λ·θ を加算してから適応的学習率で更新。結果として Weight Decay の実効値がパラメータごとの学習率に依存し、不均一になる
  • AdamW: 勾配更新と Weight Decay を分離。θ ← θ - η·(m̂/(√v̂ + ε)) - η·λ·θ。Weight Decay λ はすべてのパラメータに均一に適用される

この分離により、AdamW は以下の利点を得る:

  • 学習率と Weight Decay を独立にチューニング可能
  • 正則化の効果がパラメータの勾配スケールに左右されない
  • 大規模モデルでの過学習抑制が安定

主要 LLM での AdamW 設定値

モデル学習率β1β2Weight DecayWarmup Stepsスケジューラ
GPT-4 (推定)3e-40.90.950.12,000Cosine
Llama 3-8B3e-40.90.950.12,000Cosine
Llama 3-70B1.5e-40.90.950.12,000Cosine
Mistral Large3e-40.90.950.11,000Cosine
Qwen 2.5-72B1e-40.90.950.12,000Cosine
Gemma 2-27B1e-40.90.950.011,000Cosine
Phi-3-14B3e-40.90.950.12,000Cosine

注目すべきは β1=0.9, β2=0.95, wd=0.1 が 2024-2026 年の LLM で事実上の標準設定になっている点。β2=0.999(Adam のデフォルト)からの変更は、LLM の学習安定性を大幅に改善する。

PyTorch 実装と DeepSpeed 統合

PyTorch 2.x の標準 AdamW と DeepSpeed FusedAdam の比較:

# PyTorch 標準
from torch.optim import AdamW
optimizer = AdamW(model.parameters(), lr=3e-4, betas=(0.9, 0.95), weight_decay=0.1)

# DeepSpeed FusedAdam(CUDA カーネル融合で 15-30% 高速)
from deepspeed.ops.adam import FusedAdam
optimizer = FusedAdam(model.parameters(), lr=3e-4, betas=(0.9, 0.95), weight_decay=0.1)

DeepSpeed ZeRO Stage 2/3 環境では FusedAdam が推奨される。オプティマイザ状態を GPU 間でシャーディングし、1B パラメータモデルで GPU あたりのメモリを約 8 GB → 1-2 GB に削減できる。

メモリ使用量の詳細

AdamW は各パラメータに対して 2 つの状態を保持する:

  • 一次モーメント m(勾配の指数移動平均): FP32 で 4 bytes/param
  • 二次モーメント v(勾配二乗の指数移動平均): FP32 で 4 bytes/param
  • パラメータ本体(Master Weights): FP32 で 4 bytes/param

合計で 12 bytes/param。7B パラメータモデルなら約 84 GB のオプティマイザ+パラメータ状態が必要。BF16 混合精度学習では Forward/Backward のアクティベーションは BF16 だが、オプティマイザ状態は FP32 を維持する必要がある。

モデルサイズパラメータ (FP32)オプティマイザ状態合計
1B4 GB8 GB12 GB
7B28 GB56 GB84 GB
13B52 GB104 GB156 GB
70B280 GB560 GB840 GB

ファインチューニングでの AdamW 設定

SFT・DPO・RLHF では事前学習と異なる設定が推奨される:

  • SFT(教師ありファインチューニング): lr=1e-5〜5e-5, wd=0.01, epochs=2-3
  • LoRA/QLoRA: lr=1e-4〜3e-4, wd=0.01, LoRA rank=8-64
  • DPO: lr=5e-7〜1e-6, β2=0.999, wd=0.0(正則化は DPO の β パラメータで制御)
  • RLHF PPO: Actor lr=1e-6, Critic lr=5e-6, wd=0.01

よくある質問(FAQ)

Q: AdamW の Weight Decay はどの値に設定すべきですか?

A: LLM の事前学習では 0.1 が標準値で、Llama 3・GPT-4・Mistral Large すべてこの値を採用している。ファインチューニングでは 0.01 に下げるのが一般的。Weight Decay を 0 にすると過学習リスクが上がり、0.3 以上にすると学習が不安定になる傾向がある。

Q: β2=0.95 と β2=0.999 のどちらを使うべきですか?

A: LLM の事前学習では β2=0.95 が 2024 年以降の標準。β2=0.999(Adam のデフォルト)は二次モーメントの更新が遅すぎ、学習率の適応が追いつかない。β2=0.95 にすることで勾配の急変に素早く対応でき、Loss スパイク(突発的な損失増大)の発生が 30-50% 減少するとの報告がある。

Q: 8-bit AdamW(bitsandbytes)は品質に影響しますか?

A: bitsandbytes の 8-bit AdamW はオプティマイザ状態を INT8 に量子化し、メモリを約 75% 削減する。LoRA ファインチューニングでは FP32 AdamW とほぼ同等の品質を維持できるが、事前学習や Full Fine-tuning では 0.5-1.0% の品質低下が報告されている。メモリ制約が厳しい場合の実用的な選択肢で、QLoRA + 8-bit AdamW の組み合わせは A100 40GB 1 枚で 70B モデルの LoRA ファインチューニングを可能にする。

まとめ

  • AdamW は Weight Decay を勾配更新から分離し、安定した正則化を実現
  • β1=0.9, β2=0.95, wd=0.1, Cosine Decay が 2026 年の LLM 標準設定
  • DeepSpeed FusedAdam + ZeRO で分散学習時のメモリ・速度を最適化
  • ファインチューニングでは学習率を 1/10〜1/100 に下げ、wd=0.01 が推奨
  • 8-bit AdamW はメモリ制約下での実用的な代替手段