ひとことで言うと
正規化層は、層に入る値の尺度をそろえて学習を進めやすくする層です。バッチ正規化(BN)は同じミニバッチ内の例をまたいで統計を取り、層正規化(LN)は1例の層内のユニットをまたいで統計を取ります。名前の違いは実装の細部ではなく、正規化する軸の違いです。
同じクラスの答案をまとめて採点して基準を決めるのがBN、1人の答案の中だけで得点の尺度をそろえるのがLNです。前者はクラスの人数が少ないと基準が揺れ、後者は他の答案を待たずに採点できます。
なぜ必要か
重みの更新で前段の出力分布が変わると、後段は毎回異なる尺度の入力に適応し直します。BN論文はこの現象を internal covariate shift と呼び、層入力を正規化することで学習を速める方法を提案しました。正規化して終わりではなく、後段が必要なら元の尺度へ戻せるよう、スケールとシフトも学習させます。
ただしBNはミニバッチの情報を使うため、学習時と推論時の計算が同じではありません。学習時はそのバッチの統計、推論時は学習中に蓄積した移動平均を使います。バッチを小さくすると平均・分散の推定値が揺れ、各更新で正規化の基準まで変わります。そのため、メモリ制約で小バッチになるモデルではBNが不安定になりやすい、というのが使い分けの出発点です。
| 困りごと | 正規化層がそろえるもの | 残る判断 |
|---|---|---|
| 層入力の尺度が更新ごとに動く | 平均と分散を基準にした尺度 | スケール・シフトを学習させる |
| BNのバッチが小さい | ミニバッチ統計の揺れ | LNなど別の統計軸を検討する |
仕組み
活性化 を、ミニバッチの例 とユニット(チャネル) の値とします。BNは各 について、バッチ内の平均 と分散 を計算します。 はミニバッチサイズ、 はゼロ除算を避ける小さな定数です。
は正規化後の値、 は次の層へ渡す値です。 は学習可能なスケール、 は学習可能なシフトで、正規化が表現力を奪うならネットワーク自身が調整できます。畳み込みでは通常、空間位置も同じチャネルの統計に含めますが、軸を取り違えないことが先です。
推論時に未知のバッチの統計を使うと、同じ入力でも同居する例で出力が変わります。そこで学習中に得たバッチ統計を移動平均として更新し、推論時はその蓄積値(移動平均・移動分散)を固定して使います。ここを忘れて train と eval の挙動が変わるのが、BNの実装で最も見落としやすい点です。
LNでは、1つの例の層内ユニット から平均 と分散 を作ります。 は正規化対象の隠れ状態の次元、 はその要素です。
つまりBNは「同じユニットをバッチ方向に」、LNは「同じ例を特徴方向に」正規化します。LNは他の例も移動平均も必要としないため、学習時とテスト時が同じ計算になります。論文が述べるように、時系列モデルでは各時刻の隠れ状態ごとに統計を取りやすく、ミニバッチサイズに依存せず隠れ状態のダイナミクスを安定させる方向に働きます。
試験でどう問われるか
| 問われ方 | 正解に寄る条件 | 引っかけ |
|---|---|---|
| BNの平均・分散の軸 | 同じユニットについてミニバッチ方向に計算する | 1サンプル内の特徴方向に取る |
| BNの学習・推論差 | 学習時はバッチ統計、推論時は蓄積した移動平均・分散 | 推論時も現在のバッチ統計を使う |
| の役割 | 正規化後のスケールとシフトを学習する | 固定のハイパーパラメータとする |
| LNの軸と利点 | 1例の層内で計算し、時系列の各時刻にも適用しやすい | バッチ全体の統計が必要だとする |
| 小バッチでの判断 | BNの統計推定が揺れやすく、サイズ依存の問題になる | 小バッチほどBNの統計が正確になる |
実装で確かめる
次のコードは、同じ2次元配列にBNとLNを適用します。axis=0 のBNは列ごと、axis=1 のLNは行ごとに平均・分散を取ります。
import numpy as np
x = np.array([[1., 2., 3.], [3., 4., 5.]])
eps = 1e-5
gamma = np.ones(3)
beta = np.zeros(3)
def norm(x, axis, gamma, beta):
mean = x.mean(axis=axis, keepdims=True)
var = x.var(axis=axis, keepdims=True)
return gamma * (x - mean) / np.sqrt(var + eps) + beta
bn = norm(x, axis=0, gamma=gamma, beta=beta)
ln = norm(x, axis=1, gamma=np.ones(3), beta=np.zeros(3))
assert np.allclose(bn.mean(axis=0), 0, atol=1e-4)
assert np.allclose(ln.mean(axis=1), 0, atol=1e-4)
print("BN", bn.round(3))
print("LN", ln.round(3))
BNでは各列の平均が0、LNでは各行の平均が0になることを確認できます。ここで axis を逆にすると、式は動いても別の正規化層になります。
学習時のBN統計をそのまま推論コードへ持ち込まないでください。逆に、推論時の移動統計を学習中から固定すると、ミニバッチ統計を使うBNの定義から外れます。フレームワークのモード切替は、この統計の切替を含んでいます。
取り違えやすいもの
| 観点 | バッチ正規化(BN) | 層正規化(LN) |
|---|---|---|
| 統計を取る単位 | 同じユニットのミニバッチ内 | 1例の層内ユニット |
| 学習時 | 現在のミニバッチ統計 | その例の統計 |
| 推論時 | 移動平均・移動分散 | 学習時と同じ計算 |
| 小バッチとの相性 | 統計が揺れやすい | バッチサイズに依存しない |
| 系列モデルでの判断 | 時刻・系列の扱いを設計する必要がある | 各時刻の隠れ状態へ適用しやすい |
両者は「正規化するから同じ」ではありません。画像のように十分な例をまとめやすい場合はBNのバッチ統計が使えますが、系列モデルや小バッチではLNの軸設計が自然です。最終的には、テンソルのどの軸を統計から消すかを、入力形状とバッチサイズに照らして決めます。
想起チェック
BNとLNは、どの軸の統計を使うか
BNは同じユニットについてミニバッチ方向、LNは1例の層内ユニット方向です。
BNで学習時と推論時の統計が異なる理由は
学習時は現在のミニバッチの平均・分散を使い、推論時は蓄積した移動平均・移動分散を使うためです。推論時に他の例へ依存させないためです。
正規化後にガンマとベータを置く意味は
は学習可能なスケール、 は学習可能なシフトです。必要ならネットワークが正規化前に近い尺度を再現できます。
小バッチの系列モデルでLNを候補にする理由は
LNは他のサンプルの統計や移動平均を必要とせず、各時刻の隠れ状態に同じ計算を適用できるからです。