RAdam とは、Adam が更新の正規化に使う 2 次モーメント $v_t$ の推定が、訓練初期に大きくばらつく (分散が大きい) という分析に基づき、$v_t$ による正規化の分散を一定に保つ補正項を掛け、分散が大きすぎる最初の数ステップについては正規化自体をやめるように変更した最適化手法です。warmup が経験的に必要とされてきた Transformer の訓練に限らず、Adam による訓練全般を warmup なしに安定化させることを狙います。2020年の ICLR 2020 に採択された論文 On the Variance of the Adaptive Learning Rate and Beyond で提案されました。
この研究の位置付け (リンク先はこれまでに書いた文献メモ)
- [Saxe2014] 深層線形ニューラルネットの学習ダイナミクスの厳密解と直交初期化 — 初期化・学習の安定性を理論的に扱った源流的な研究。
- [Zhang2019] RMSNorm (LayerNorm から中心化を省いた正規化) — [Xiong2020] と組む Pre-RMSNorm として、現代 LLM の正規化で主流。
- [Liu2020] RAdam (本記事) — Adam の 2 次モーメントの分散を補正し、最適化側から warmup を不要化。初期化側の [Huang2020] と対をなすが、現代 LLM では [Xiong2020] の Pre-LN と warmup 付き Adam の併用が一般的。
- [Xiong2020] Pre-LN — 正規化を残差の中に置いて warmup なしで安定化。warmup 問題への「正規化層の配置」からの解で、現代 LLM の主流。
- [Huang2020] T-Fixup — 初期化を小さくして LayerNorm と warmup を不要化。warmup 問題への「初期化」からの解で、最適化側の [Liu2020] と同じ問題意識。
- [Stollenwerk2025] Coupled Adam — Adam の 2 次モーメントが埋込みを偏らせる問題への対処。2 次モーメントに着目する点で [Liu2020] と関連する近年の研究。
以下に文献概要とトイコードを記します。
参考文献
Liyuan Liu, Haoming Jiang, Pengcheng He, Weizhu Chen, Xiaodong Liu, Jianfeng Gao, and Jiawei Han. On the Variance of the Adaptive Learning Rate and Beyond. In Proceedings of the 8th International Conference on Learning Representations (ICLR 2020), 2020.
文献概要
$1/\sqrt{v_t}$ の分散が $\mathrm{Var} \propto \rho_t/((\rho_t-2)(\rho_t-4))$ となる部分 ($v_t$ を自由度 $\rho_t$ のカイ二乗分布とみなして得る分散の近似) だけは原論文に基づき、私は追っていません。$\rho_t$ と $r_t$ の式は、これを認めれば本文のとおり導けます。
- Adam (私の検算記事) は、1, 2 次モーメントを用いる勾配降下法です。更新は、1 次モーメント $m_t$ (平均的な勾配方向) を、2 次モーメント $v_t$ (平均的な勾配の大きさ) で正規化する形になっています (ただし訓練初期の補正と安定化の $\varepsilon$ を省略)。
- $w_t = w_{t-1} - \alpha \cdot ( m_t / \sqrt{v_t})$ ($\alpha$ は学習率)
- ところで、オリジナルの Transformer などの訓練には学習率の warmup (訓練初期の学習率を小さく始めて徐々に上げる) が経験的に必要でした (T-Fixup の記事 参照)。Liu らは、warmup がなぜ効くのかを、Adam が正規化に使う $v_t$ の推定の分散 という観点から分析しました。
- 訓練初期は、2 次モーメント $v_t$ の推定に使えたサンプルが少ないため、$v_t$ の推定の分散が非常に大きくなります (極端には最初の数ステップで発散しうる) [1]。この分散の大きい $v_t$ で正規化するので、訓練初期の更新の分散も大きくなり、学習が不安定になります。
- warmup は初期の学習率を下げることで、この分散の影響を抑えている、というのが Liu らの解釈です [1]。
- そこで Liu らは、この $v_t$ による正規化の分散を解析的に見積もり、それを一定に保つことを考えました。分散を「$v_t$ が実質いくつの $g^2$ を平均したものか」で表したいところ、$v_t$ は重みの減衰する移動平均なので、まず実効サンプル数 $\rho_t$ を求めておきます(Adam は通常 1 に近い減衰率で使うのでこれはほぼステップ数 $t$ に近いのですが、一般の減衰率でも正しく扱うため、分散と補正項 $r_t$ を $\rho_t$ の式で表すため、分散が有限になる範囲($\rho_t > 4$)を判定するために求めておきます)。
-
$\rho_t$ の求め方 (重みの重心を合わせる) : 減衰付き移動平均で $k$ ステップ前の $g^2$ にかかる重みは $\beta_2^k$ に比例します。この重みの重心 (平均の遅れ) を、長さ $\rho$ の減衰なし窓 ($0\sim\rho-1$ ステップ前を等しい重み、重心は $(\rho-1)/2$) の重心に一致させて $\rho$ を決めます。
- 定常状態 ($t\to\infty$) の重心は $\sum_{k\ge0} k\beta_2^k \big/ \sum_{k\ge0}\beta_2^k = \beta_2/(1-\beta_2)$ です。$(\rho-1)/2 = \beta_2/(1-\beta_2)$ を解くと、最大長 $\rho_\infty = 2/(1-\beta_2) - 1$ を得ます。
- 訓練初期は過去 $t$ 個しかなく重心が小さいので、$\sum_{k=0}^{t-1} k\beta_2^k \big/ \sum_{k=0}^{t-1}\beta_2^k$ を $(\rho_t-1)/2$ に一致させ、等比級数を整理すると $\rho_t = \rho_\infty - 2t,\beta_2^t/(1-\beta_2^t)$ となります ($\rho_t$ は $t$ とともに $\rho_\infty$ へ増えます)。
- $\rho_t$ を求めると何が嬉しいか : $v_t$ を「$\rho_t$ 個のサンプルの平均」とみなせると、$1/\sqrt{v_t}$ の分散が $\rho_t$ だけの関数で書けます。$\rho_t$ が大きいほど分散は小さく、$\rho_t \le 4$ では発散します(これが手続きで $\rho_t \le 4$ のとき正規化をやめる理由です)。
-
補正項 $r_t$ の決め方 : 分散を定常状態 ($\rho_\infty$) の値に合わせるように $r_t$ を決めます。$r_t$ は $t$ が進むと 1 に近づきます。
- 原論文の近似では $\mathrm{Var}(1/\sqrt{v_t}) \propto \rho_t/((\rho_t-2)(\rho_t-4))$ で、これを定常値に合わせると ($r_t^2 = \mathrm{Var}_{\rho_\infty}/\mathrm{Var}_{\rho_t}$) 以下となります。分母の $(\rho_t-4)$ から、分散が有限なのは $\rho_t>4$ に限られるとわかります。
- $r_t = \sqrt{\dfrac{(\rho_t - 4)(\rho_t - 2), \rho_\infty}{(\rho_\infty - 4)(\rho_\infty - 2), \rho_t}}$
- 原論文の近似では $\mathrm{Var}(1/\sqrt{v_t}) \propto \rho_t/((\rho_t-2)(\rho_t-4))$ で、これを定常値に合わせると ($r_t^2 = \mathrm{Var}_{\rho_\infty}/\mathrm{Var}_{\rho_t}$) 以下となります。分母の $(\rho_t-4)$ から、分散が有限なのは $\rho_t>4$ に限られるとわかります。
-
$\rho_t$ の求め方 (重みの重心を合わせる) : 減衰付き移動平均で $k$ ステップ前の $g^2$ にかかる重みは $\beta_2^k$ に比例します。この重みの重心 (平均の遅れ) を、長さ $\rho$ の減衰なし窓 ($0\sim\rho-1$ ステップ前を等しい重み、重心は $(\rho-1)/2$) の重心に一致させて $\rho$ を決めます。
-
RAdam の手続き : 各ステップ $t$ で、Adam と同じく 1 次・2 次モーメント $m_t, v_t$ を計算したうえで、$\rho_t$ に応じて次のようにパラメータ $w$ を更新します ($\alpha$ は学習率。簡単のためモーメントのバイアス補正と $\varepsilon$ は省略)。
- 実効サンプル数を求める : $\rho_t = \rho_\infty - 2t,\beta_2^t/(1-\beta_2^t)$ ($\rho_\infty = 2/(1-\beta_2)-1$)。
- $\rho_t \le 4$ のとき : $v_t$ による正規化をやめ、モメンタム付き SGD で更新する。$w_t = w_{t-1} - \alpha, m_t$。
-
$\rho_t > 4$ のとき : 補正項 $r_t$ を掛けて、Adam と同じく $v_t$ で正規化して更新する。$w_t = w_{t-1} - \alpha, r_t, m_t/\sqrt{v_t}$。
- $\rho_t$ が 4 を超えた直後は補正項 $r_t$ が小さいので、更新幅はかなり小さくなります (トイコード参照)。
- 原論文の検証実験では、画像分類・言語モデリング・ニューラル機械翻訳などで、warmup を用いた Adam と同等の性能を warmup の調整なしに得られること、および学習率の選び方に頑健になることが報告されています [1]。
現在の状況
RAdam は PyTorch に torch.optim.RAdam として組み込まれています。ただし標準として広く定着したわけではなく、大規模訓練では Adam / AdamW + warmup が一般的なようです。
トイコード
このトイコードは手続きを追えるように RAdam を自前実装しています。
以下は、RAdam を自前実装し、$\beta_2 = 0.999$ のときの補正スケジュールを表示するトイコードです。$\rho_t$ と $r_t$ は勾配に依存せず $t$ と $\beta_2$ だけで決まるので、ここではダミー勾配で数ステップ回して、各 $t$ での値を取り出しています。
import torch
class RAdam(torch.optim.Optimizer):
# On the Variance of the Adaptive Learning Rate and Beyond (ICLR 2020)
def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8):
defaults = {'lr': lr, 'betas': betas, 'eps': eps}
super().__init__(params, defaults)
@torch.no_grad()
def step(self):
for group in self.param_groups:
lr = group['lr']
beta1, beta2 = group['betas']
eps = group['eps']
rho_inf = 2 / (1 - beta2) - 1
for p in group['params']:
if p.grad is None:
continue
grad = p.grad
state = self.state[p]
if len(state) == 0:
state['step'] = 0
state['exp_avg'] = torch.zeros_like(p)
state['exp_avg_sq'] = torch.zeros_like(p)
state['step'] += 1
t = state['step']
exp_avg = state['exp_avg']
exp_avg_sq = state['exp_avg_sq']
exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
m_hat = exp_avg / (1 - beta1 ** t)
# 2 次モーメントの減衰付き移動平均を「長さ rho_t の減衰なし移動平均」とみなす
rho_t = rho_inf - 2 * t * beta2 ** t / (1 - beta2 ** t)
if rho_t > 4: # v_t の推定が安定 -> v_t で正規化し補正項 r_t を掛ける
v_hat = (exp_avg_sq / (1 - beta2 ** t)).sqrt().add_(eps)
r_t = (((rho_t - 4) * (rho_t - 2) * rho_inf) /
((rho_inf - 4) * (rho_inf - 2) * rho_t)) ** 0.5
p.addcdiv_(m_hat, v_hat, value=-lr * r_t)
state['rho_t'], state['r_t'] = rho_t, r_t
else: # v_t の推定が不安定 -> 正規化せずモメンタム付き SGD
p.add_(m_hat, alpha=-lr)
state['rho_t'], state['r_t'] = rho_t, None
def main():
param = torch.nn.Parameter(torch.zeros(1))
opt = RAdam([param], lr=0.1, betas=(0.9, 0.999))
print('beta2 = 0.999, rho_inf = 1999')
print(f'{"t":>5} | {"rho_t":>9} | {"update":<12} | {"r_t":>8}')
for t in [1, 2, 3, 4, 5, 10, 100, 1000]:
while opt.state.get(param, {}).get('step', 0) < t:
param.grad = torch.ones(1) # ダミー勾配 (rho_t, r_t は勾配に依存しない)
opt.step()
st = opt.state[param]
if st['r_t'] is None:
print(f'{t:>5} | {st["rho_t"]:>9.3f} | {"momentum SGD":<12} | {"-":>8}')
else:
print(f'{t:>5} | {st["rho_t"]:>9.3f} | {"Adam * r_t":<12} | {st["r_t"]:>8.4f}')
if __name__ == '__main__':
main()
実行すると以下のようになります。
beta2 = 0.999, rho_inf = 1999
t | rho_t | update | r_t
1 | 1.000 | momentum SGD | -
2 | 1.999 | momentum SGD | -
3 | 2.999 | momentum SGD | -
4 | 3.997 | momentum SGD | -
5 | 4.996 | Adam * r_t | 0.0173
10 | 9.983 | Adam * r_t | 0.0490
100 | 98.333 | Adam * r_t | 0.2153
1000 | 835.967 | Adam * r_t | 0.6453
読み取れることは以下です。
- 最初の 4 ステップ ($t \leq 4$) : $\rho_t \leq 4$ なので $v_t$ による正規化をやめ、モメンタム付き SGD になります。$v_t$ の推定サンプルが少なく、その分散が大きすぎるためです。
- $t = 5$ 以降 : $v_t$ による正規化を再開しますが、補正項 $r_t$ は $0.0173$ と非常に小さく、そこから $t$ とともに $1$ へ徐々に増えていきます。分散を一定に保とうとした結果として、学習率の warmup を明示的に設定しなくても、訓練初期ほど更新が抑えられます。
- T-Fixup の記事 は初期化を工夫して warmup を外す話でした。本記事は最適化側から warmup を外す話です。どちらも「Adam の訓練初期の大きな更新」という同じ問題への、別方向のアプローチといえます。