0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

初期化直後の Transformer における中間出力のノルムのスケール

0
Last updated at Posted at 2026-07-25

Transformer は Pre-LN 構成が安定すると主張した研究 [1] [関連記事] では、初期化直後における最終層の重みの勾配のオーダーを見積もっています。その準備として、初期化直後のエンコーダ層内の各中間出力のノルムのスケールを評価する補題 2が使われます (Pre-LN 構成では層を経るごとにスケールが大きくなる)。

この記事では補題 2と、その証明のなかで使う補題 1を追います。なお、補題 2の前に、補題 2で用いる表記法を導入します。

参考文献

  1. 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$ となります。

0
0
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?