Transformer は Pre-LN 構成が安定すると主張した研究 [1] [関連記事1] では、LayerNorm の入力に関するヤコビ行列のスペクトルノルムのオーダー $\|\mathbf{J}_{LN}(x)\|_{2}=\mathcal{O}(\sqrt{d} / \|x\|_{2})$ は入力変数のノルムが大きいほど小さくなるが、Pre-LN 構成では残差接続後に正規化しないために入力変数のノルムがエンコーダ層ごとに積み増しされていくので $\|\mathbf{J}_{LN}(x)\|_{2}$ が抑えられる、といった議論をしています。
この記事では、なぜオーダーがそうなるかの補題だけ追います。原論文の補題 3がそれですが、そこで必要な補題 4.1も先に書きます (補題 4.1 は Appendix にだけ登場するのでナンバリングが前後しています)。
参考文献
- 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$)。
補題 4.1
$\alpha\in\mathbb{R}^{d}$ を $\|\alpha\|_{2}=1$ なるベクトルとすると、$I-\alpha^{\top}\alpha$ の固有値は 1 または 0 である。
補題 4.1 の証明
$e_1=\alpha$ を含む正規直交基底 ${e_1,\dots,e_d}$ をとると、
\begin{align}
e_1(I-\alpha^{\top}\alpha)&=e_1-e_1\alpha^{\top}\alpha=e_1-\alpha=0 \\
e_i(I-\alpha^{\top}\alpha)&=e_i-e_i\alpha^{\top}\alpha=e_i \quad (i\neq 1)
\end{align}
となる。よって ${e_1,\dots,e_d}$ は $I-\alpha^{\top}\alpha$ の固有ベクトルであり、対応する固有値は $(0,1,1,\dots,1)$ である。
補題 3
$x\in\mathbb{R}^{d}$ とする。$\text{LN}(x)$ を、スケール $\gamma=1, $ バイアス $\beta=0$ に初期化した LayerNorm とする。この初期状態で、 $\|\mathbf{J}_{LN}(x)\|_{2}=\mathcal{O}(\dfrac{\sqrt{d}}{\|x\|_{2}})$ が成り立つ。ここで $\mathbf{J}_{LN}(x)=\dfrac{\partial\text{LN}(x)}{\partial x}$ は $\text{LN}(x)$ のヤコビ行列である。
補題 3 の証明
$\mathbf{1}=(1,1,\dots,1)\in\mathbb{R}^{d}$ とし、$y=x(I-\dfrac{1}{d}\mathbf{1}^{\top}\mathbf{1})$ とおくと、$y$ は $x$ から平均を引いたものになるので (補足1)、LayerNorm は以下のように書ける。
\text{LN}(x)_{i}=\frac{y_{i}}{\sqrt{\frac{1}{d}\sum_{j=1}^{d}y_{j}^{2}}}
この偏微分は、商の微分により
\begin{align}
\frac{\partial\text{LN}(x)_{i}}{\partial y_{j}}
&=\frac{\partial}{\partial y_{j}}\left(\frac{y_{i}}{\sqrt{\frac{1}{d}\sum_{k=1}^{d}y_{k}^{2}}}\right) \\
&=\frac{\delta_{ij}\sqrt{\frac{1}{d}\sum_{k=1}^{d}y_{k}^{2}}-y_{i}\dfrac{\frac{1}{d}y_{j}}{\sqrt{\frac{1}{d}\sum_{k=1}^{d}y_{k}^{2}}}}{\frac{1}{d}\sum_{k=1}^{d}y_{k}^{2}} \\
&=\frac{\sqrt{d}}{\|y\|_{2}}\left(\delta_{ij}-\frac{y_{i}y_{j}}{\|y\|_{2}^{2}}\right)
\end{align}
となる。ここで $\delta_{ij}$ は $i=j$ のとき $1$、$i\neq j$ のとき $0$ である。これを行列にまとめると
\frac{\partial\text{LN}(x)}{\partial y}=\frac{\sqrt{d}}{\|y\|_{2}}\left(I-\frac{y^{\top}y}{\|y\|_{2}^{2}}\right)
となる。さらに $y=x(I-\dfrac{1}{d}\mathbf{1}^{\top}\mathbf{1})$ より $\dfrac{\partial y}{\partial x}=I-\dfrac{1}{d}\mathbf{1}^{\top}\mathbf{1}$ なので、
\begin{align}
\mathbf{J}_{LN}(x)=\frac{\partial\text{LN}(x)}{\partial x}
&=\frac{\partial\text{LN}(x)}{\partial y}\frac{\partial y}{\partial x} \\
&=\sqrt{d}\,\frac{1}{\|y\|_{2}}\left(I-\frac{y^{\top}y}{\|y\|_{2}^{2}}\right)\left(I-\frac{1}{d}\mathbf{1}^{\top}\mathbf{1}\right)
\end{align}
となる。ここで、補題 4.1 より $\left(I-\dfrac{y^{\top}y}{\|y\|_{2}^{2}}\right)$ と $\left(I-\dfrac{1}{d}\mathbf{1}^{\top}\mathbf{1}\right)$ の固有値は 1 または 0 なので $\Biggl\|(I-\dfrac{y^{\top}y}{\|y\|_{2}^{2}})\Biggr\|_{2}=1$ かつ $\Biggl\|(I-\dfrac{1}{d}\mathbf{1}^{\top}\mathbf{1})\Biggr\|_{2}=1$ であり (補足2)、以下となる。
\|\mathbf{J}_{LN}(x)\|_{2}=\mathcal{O}\left(\frac{\sqrt{d}}{\|y\|_{2}}\right)=\mathcal{O}\left(\frac{\sqrt{d}}{\|x\|_{2}}\right)
補足を書きます。
補足 1
$\mathbf{1}^{\top}\mathbf{1}$ はすべての成分が1の行列です。これを $x$ に右からかけると
$$
x(\mathbf{1}^{\top}\mathbf{1})=\left(\sum_{k=1}^{d}x_k,\ \sum_{k=1}^{d}x_k,\ \dots,\ \sum_{k=1}^{d}x_k\right)
$$
となります。これを $d$ で割ると
$$
x\left(\frac{1}{d}\mathbf{1}^{\top}\mathbf{1}\right)=\left(\bar{x},\bar{x},\dots,\bar{x}\right)
$$
となり、全成分が $\bar{x}$ であるベクトルになります。よって、
$$
y=x-x\left(\frac{1}{d}\mathbf{1}^{\top}\mathbf{1}\right)
$$
は、各成分が $y_i=x_i-\bar{x}$ となり、$x$ の各成分から平均を引いたものになります。
補足 2
実対称行列のスペクトルノルムは、絶対値が最大の固有値の絶対値に等しいです。