Transformer は Pre-LN 構成が安定すると主張した研究 [1] [関連記事] では、初期化直後における最終層の重みの勾配のオーダーを見積もっています。その準備として、初期化直後のエンコーダ層内の各中間出力のノルムのスケールを評価する補題 2が使われます (Pre-LN 構成では層を経るごとにスケールが大きくなる)。
この記事では補題 2と、その証明のなかで使う補題 1を追います。なお、補題 2の前に、補題 2で用いる表記法を導入します。
参考文献
- Ruibin Xiong, Yunchang Yang, Di He, Kai Zheng, Shuxin Zheng, Chen Xing, Huishuai Zhang, Yanyan Lan, Liwei Wang, and Tie-Yan Liu. On Layer Normalization in the Transformer Architecture. Proceedings of the 37th International Conference on Machine Learning (ICML 2020), 2020.
https://proceedings.mlr.press/v119/xiong20b.html
関連記事:
1. 【文献メモ】現代 LLM の正規化 Pre-LN + RMSNorm の成り立ち
2. LayerNorm の入力に関するヤコビ行列のスペクトルノルムのオーダー
3. 初期化直後の Transformer における中間出力のノルムのスケール (これ)
以下、ベクトルは横ベクトルとします (入力を重みに左から掛ける流儀 $y=xW$)。また、パラメータは文献 [1] にしたがい以下のように初期化された直後を考えます。
- 各線形変換の重み行列は Xavier 初期化する。すなわち $n_{\rm in}\times n_{\rm out}$ 行列の各要素を独立に $N(0, \dfrac{2}{n_{\rm in}+n_{\rm out}})$ からサンプリングする。特に次元 $d$ の正方行列では $N(0, 1/d)$ となる。
- 各線形変換のバイアスはゼロベクトルで初期化する。
- LayerNorm のスケールとバイアスは $\gamma=1,$ $\beta=0$ で初期化する。このとき $\text{LN}(v)=\dfrac{v-\mu}{\sigma}$ となる ($\mu,$ $\sigma$ は $v$ の成分の平均・標準偏差)。
- マルチヘッドセルフアテンション (MHA) 層の $W^{Q,l}, W^{K,l}$ はゼロ行列で初期化する。よって初期化直後のアテンションは一様分布となり、MHA 出力は各トークンの値の平均を $W^{V,l}$ で線形変換したものになる。
- 入力ベクトルの各成分も、重み行列 ($d\times d$ の場合) と同じ $N(0, 1/d)$ からサンプリングされているとする (入力はトークン埋め込みと位置埋め込みの線形結合だが、いずれもガウス分布で初期化するイメージ)。
補題 1
確率ベクトル $X\in\mathbb{R}^{d}$ が $N(0,\sigma^{2}\mathbf{I}_{d})$ にしたがうならば、$\mathbb{E}(\|\text{ReLU}(X)\|_{2}^{2})=\dfrac{1}{2}\sigma^{2}d$ である。
補題 1 の証明
各成分が独立に $N(0,\sigma^{2})$ にしたがう $X=(X_1,\dots,X_d)$ を考える。$X_1$ の確率密度関数を $\rho_X$ とすると、以下となる。
\begin{align}
\mathbb{E}(\|\text{ReLU}(X)\|_{2}^{2})
&=\sum_{i=1}^{d}\mathbb{E}(\text{ReLU}(X_i)^{2}) \\
&=\sum_{i=1}^{d}\mathbb{E}(\text{ReLU}(X_i)^{2}\mid X_i\geqq 0)\,\mathbb{P}(X_i\geqq 0) \\
&=\frac{d}{2}\,\mathbb{E}(X_1^{2}\mid X_1\geqq 0) \\
&=\frac{d}{2}\int_{0}^{\infty}x^{2}\,2\rho_X(x)\,dx \\
&=\frac{1}{2}\sigma^{2}d
\end{align}
補題 2 で用いる表記法
初期化直後の $l$ 番目のエンコーダ層の内部で、第 $i$ トークンがたどる中間出力を次のように書く。
Post-LN Transformer では、入力 $x^{{\rm post}}_{l,i}$ から順に以下をたどる。
- $x^{{\rm post},1}_{l,i}$:MHA 後
- $x^{{\rm post},2}_{l,i}$:1 回目の残差接続後
- $x^{{\rm post},3}_{l,i}$:1 回目の LayerNorm 後
- $x^{{\rm post},4}_{l,i}$:FFN 後
- $x^{{\rm post},5}_{l,i}$:2 回目の残差接続後
- $x^{{\rm post}}_{l+1,i}$:2 回目の LayerNorm 後 (次のエンコーダ層への入力)
Pre-LN Transformer では、入力 $x^{{\rm pre}}_{l,i}$ から順に以下をたどる。
- $x^{{\rm pre},1}_{l,i}$:1 回目の LayerNorm 後
- $x^{{\rm pre},2}_{l,i}$:MHA 後
- $x^{{\rm pre},3}_{l,i}$:1 回目の残差接続後
- $x^{{\rm pre},4}_{l,i}$:2 回目の LayerNorm 後
- $x^{{\rm pre},5}_{l,i}$:FFN 後
- $x^{{\rm pre}}_{l+1,i}$:2 回目の残差接続後 (次のエンコーダ層への入力)
その他の記号は次のとおり。
- $\text{FFN}(x)=\text{ReLU}(xW^{1,l})W^{2,l}$:FFN の定義
- $n$:1 系列あたりのトークン数
- $d$:トークンの埋め込み次元数
補題 2
初期化直後の Post-LN Transformer は任意の $l, i$ で $\mathbb{E}(\|x_{l,i}^{{\rm post},5}\|_{2}^{2})=\dfrac{3}{2}d$ を満たす。また、初期化直後の Pre-LN Transformer は任意の $l, i$ で $(1+\dfrac{l}{2})d\leqq\mathbb{E}(\|x^{{\rm pre}}_{l,i}\|_{2}^{2})\leqq(1+\dfrac{3l}{2})d$ を満たす。なお、ここでの期待値は入力および初期化に関する期待値である。
補題 2 の証明
初期化直後の LayerNorm は $\text{LN}(v)=\dfrac{v-\mu}{\sigma}$ であり、任意のベクトル $v$ を半径 $\sqrt{d}$ の $(d-1)$-球面上へ射影する。つまり、
\|\text{LN}(v)\|_{2}^{2}=\left\|\frac{v-\mu}{\sigma}\right\|_{2}^{2}=\frac{\sum_{k=1}^{d}(v_{k}-\mu)^{2}}{\sigma^{2}}=d
となる。以下、これを用いて各中間出力のノルムの期待値を順に評価する。
Post-LN Transformer の場合
$\|x_{l,i}^{{\rm post},3}\|_{2}^{2}$ は必ず $d$ に正規化されるので、そこからの変化を考える (層入力から $x_{l,i}^{{\rm post},3}$ までのノルムの変化については補足1を参照)。
FFN 後 $x_{l,i}^{{\rm post},4}=\text{ReLU}(x_{l,i}^{{\rm post},3}W^{1,l})W^{2,l}$ を評価すると、補題 1 を用いて
\begin{align}
\mathbb{E}(\|x_{l,i}^{{\rm post},4}\|_{2}^{2})
&=\mathbb{E}(\|\text{ReLU}(x_{l,i}^{{\rm post},3}W^{1,l})\|_{2}^{2}) \\
&=\mathbb{E}(\dfrac{1}{2}\|x_{l,i}^{{\rm post},3}\|_{2}^{2})=\frac{d}{2}
\end{align}
を得る (補足2)。これをもとに 2 回目の残差接続後 $x_{l,i}^{{\rm post},5}=x_{l,i}^{{\rm post},3}+x_{l,i}^{{\rm post},4}$ を評価すると
\begin{align}
\mathbb{E}(\|x_{l,i}^{{\rm post},5}\|_{2}^{2})
&=\mathbb{E}(\|x_{l,i}^{{\rm post},3}\|_{2}^{2})+\mathbb{E}(\|x_{l,i}^{{\rm post},4}\|_{2}^{2})+2\mathbb{E}(x_{l,i}^{{\rm post},3}{x_{l,i}^{{\rm post},4}}^{\top}) \\
&=\mathbb{E}(\|x_{l,i}^{{\rm post},3}\|_{2}^{2})+\mathbb{E}(\|x_{l,i}^{{\rm post},4}\|_{2}^{2}) \\
&=d+\frac{d}{2}=\frac{3}{2}d
\end{align}
となる。ここで 2 番目の等号で交差項の期待値が消えること (補足1) を用いた。これで Post-LN 側の主張が示された。
Pre-LN Transformer の場合
1 回目の残差接続後 $x_{l,i}^{{\rm pre},3}=x_{l,i}^{{\rm pre}}+x_{l,i}^{{\rm pre},2}$ について、MHA 出力が $x_{l,i}^{{\rm pre},2}=\dfrac{1}{n}\sum_{j=1}^{n}x_{l,j}^{{\rm pre},1}W^{V,l}$ であることと、交差項の期待値が消えること (補足1) から、
\begin{align}
\mathbb{E}(\|x_{l,i}^{{\rm pre},3}\|_{2}^{2})
&=\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})+\mathbb{E}(\|x_{l,i}^{{\rm pre},2}\|_{2}^{2}) \\
&=\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})+\mathbb{E}(\|\frac{1}{n}\sum_{i=1}^{n}x_{l,i}^{{\rm pre},1}\|_{2}^{2})
\end{align}
となる。Pre-LN では MHA の直前に LayerNorm がかかるため $\|x_{l,i}^{{\rm pre},1}\|_{2}^{2}=d$ であり、$0\leq\mathbb{E}(\|\dfrac{1}{n}\sum_{i}x_{l,i}^{{\rm pre},1}\|_{2}^{2})\leq d$ が成り立つ (補足3)。よって
\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})\leqq\mathbb{E}(\|x_{l,i}^{{\rm pre},3}\|_{2}^{2})\leqq\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})+d
を得る。次に、Pre-LN では 2 回目の残差接続後がそのまま次層の入力 $x_{l+1,i}^{{\rm pre}}=x_{l,i}^{{\rm pre},3}+x_{l,i}^{{\rm pre},5}$ になる。FFN の直前の LayerNorm から $\|x_{l,i}^{{\rm pre},4}\|_{2}^{2}=d$ であり、補題 1 により $\mathbb{E}(\|x_{l,i}^{{\rm pre},5}\|_{2}^{2})=\dfrac{1}{2}d$ となる (Post-LN の $x_{l,i}^{{\rm post},4}$ とまったく同じ計算)。よって交差項の期待値が消えること (補足1) から、
\begin{align}
\mathbb{E}(\|x_{l+1,i}^{{\rm pre}}\|_{2}^{2})
&=\mathbb{E}(\|x_{l,i}^{{\rm pre},3}\|_{2}^{2})+\mathbb{E}(\|x_{l,i}^{{\rm pre},5}\|_{2}^{2}) \\
&=\mathbb{E}(\|x_{l,i}^{{\rm pre},3}\|_{2}^{2})+\frac{1}{2}d
\end{align}
となる。これと $x_{l,i}^{{\rm pre},3}$ に関する不等式を合わせると、
\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})+\frac{1}{2}d\leqq\mathbb{E}(\|x_{l+1,i}^{{\rm pre}}\|_{2}^{2})\leqq\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})+\frac{3}{2}d
が得られる。この漸化的な不等式を層について繰り返すと、
(1+\frac{l}{2})d\leqq\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})\leqq(1+\frac{3l}{2})d
となる (補足4)。
補足を書きます。
補足 1
Post-LN Transformer のエンコーダ層を最初から追うと、MHA 出力は $x_{l,i}^{{\rm post},1}=\dfrac{1}{n}\sum_{j=1}^{n}x_{l,j}^{{\rm post}}W^{V,l}$ で、$x_{l,i}^{{\rm post},2}=x_{l,i}^{{\rm post}}+x_{l,i}^{{\rm post},1}$ なので、
\begin{align}
\mathbb{E}(\|x_{l,i}^{{\rm post},2}\|_{2}^{2})
&=\mathbb{E}(\|x_{l,i}^{{\rm post}}\|_{2}^{2})+\mathbb{E}(\|x_{l,i}^{{\rm post},1}\|_{2}^{2})+2\mathbb{E}(x_{l,i}^{{\rm post},1}{x_{l,i}^{{\rm post}}}^{\top}) \\
&=\mathbb{E}(\|x_{l,i}^{{\rm post}}\|_{2}^{2})+\mathbb{E}(\|x_{l,i}^{{\rm post},1}\|_{2}^{2}) \\
&=\mathbb{E}(\|x_{l,i}^{{\rm post}}\|_{2}^{2})+\mathbb{E}(\|\frac{1}{n}\sum_{i=1}^{n}x_{l,i}^{{\rm post}}\|_{2}^{2})
\end{align}
となります。1 番目から 2 番目の等号では、交差項 $2\mathbb{E}(x_{l,i}^{{\rm post},1}{x_{l,i}^{{\rm post}}}^{\top})$ が消えることを使いました。これは、$x_{l,i}^{{\rm post},1}$ が末尾に平均ゼロの重み $W^{V,l}$ を掛けた形をしており、$W^{V,l}$ が $x_{l,i}^{{\rm post}}$ 側と独立で期待値 $0$ だからです。一般に、残差接続 $a+b$ において $b$ が MHA や FFN の出力 (最後に平均ゼロの重み $W^{V,l}$ または $W^{2,l}$ を掛けた形) であれば、同様に交差項の期待値 $\mathbb{E}(ab^{\top})$ は $0$ になります。
なお、$x_{l,i}^{{\rm post},2}$ は続く LayerNorm によって $x_{l,i}^{{\rm post},3}$ に正規化されるため、上で求めた結果そのものは最終結果には効きません。
補足 2
FFN 後 $x_{l,i}^{{\rm post},4}=\text{ReLU}(x_{l,i}^{{\rm post},3}W^{1,l})W^{2,l}$ の 2 乗ノルムの期待値を、まず $W^{2,l}$ について、次に $W^{1,l}$ について順に取ります。行ベクトル $z$ を固定すると、$W^{2,l}$ の各要素は独立に $N(0,1/d)$ にしたがうので、
\mathbb{E}_{W^{2,l}}(\|zW^{2,l}\|_{2}^{2})=\|z\|_{2}^{2}\cdot d\cdot\frac{1}{d}=\|z\|_{2}^{2}
となります。よって $z=\text{ReLU}(x_{l,i}^{{\rm post},3}W^{1,l})$ を残して $\mathbb{E}(\|x_{l,i}^{{\rm post},4}\|_{2}^{2})=\mathbb{E}(\|\text{ReLU}(x_{l,i}^{{\rm post},3}W^{1,l})\|_{2}^{2})$ を得ます。さらに $x_{l,i}^{{\rm post},3}$ を固定すると、$x_{l,i}^{{\rm post},3}W^{1,l}$ の各成分は独立に $N(0,\|x_{l,i}^{{\rm post},3}\|_{2}^{2}/d)$ にしたがうので、$\sigma^{2}=\|x_{l,i}^{{\rm post},3}\|_{2}^{2}/d$ として補題 1を適用すると $\mathbb{E}(\|\text{ReLU}(x_{l,i}^{{\rm post},3}W^{1,l})\|_{2}^{2}\mid x_{l,i}^{{\rm post},3})=\dfrac{1}{2}\|x_{l,i}^{{\rm post},3}\|_{2}^{2}$ となります。$\|x_{l,i}^{{\rm post},3}\|_{2}^{2}=d$ を代入して $\dfrac{d}{2}$ を得ます。
補足 3
$v_i:=x_{l,i}^{{\rm pre},1}$ は各 $\|v_i\|_{2}^{2}=d$ を満たします。2 乗ノルムは凸関数なので、Jensen 不等式より
\left\|\frac{1}{n}\sum_{i=1}^{n}v_i\right\|_{2}^{2}\leqq\frac{1}{n}\sum_{i=1}^{n}\|v_i\|_{2}^{2}=d
が成り立ちます。期待値をとっても不等号は保たれます。
補足 4
$a_l:=\mathbb{E}(\|x_{l,i}^{{\rm pre}}\|_{2}^{2})$ とおくと、示した不等式は $a_l+\dfrac{1}{2}d\leqq a_{l+1}\leqq a_l+\dfrac{3}{2}d$、すなわち 1 層進むごとにノルムの 2 乗の期待値が $\dfrac{1}{2}d$ 以上 $\dfrac{3}{2}d$ 以下だけ増えることを意味します。これを入力側から積み重ねると、$l$ 層目では下限が $\dfrac{1}{2}d$ の $l$ 個ぶん、上限が $\dfrac{3}{2}d$ の $l$ 個ぶんだけ基準 $d$ に加わり、$(1+\dfrac{l}{2})d\leqq a_l\leqq(1+\dfrac{3l}{2})d$ となります。