BatchNorm vs LayerNormの理論と使い分け

深層ネットワークの学習で、ある層のパラメータが更新されると、次の層への入力の分布が変化します。すると次の層は「動く的を撃つ」ような状態になり、学習が不安定になったり遅くなったりします。この現象を内部共変量シフト(Internal Covariate Shift, ICS)と呼びます。

たとえるなら、料理のレシピを練習しているのに、毎回オーブンの温度が勝手に変わるようなものです。レシピの手順を覚えても、オーブンの温度が安定しなければ上手に焼けません。正規化層は「オーブンの温度を一定に保つ」役割を果たし、各層が安定した入力分布のもとで学習できるようにします。

2015年に提案されたBatch Normalization(BatchNorm) は、ミニバッチ内の統計量で正規化を行い、深層学習の学習速度と安定性を劇的に改善しました。その後、バッチサイズに依存しないLayer Normalization(LayerNorm) が提案され、TransformerやRNNで標準的に使われるようになりました。

正規化手法を理解すると、以下の判断ができるようになります。

  • CNNの設計: なぜBatchNormが畳み込みネットワークで効果的なのか
  • Transformerの設計: なぜLayerNormが使われるのか
  • 小バッチ学習: バッチサイズが小さいときの正規化戦略
  • 学習の安定化: 正規化が勾配の流れをどう改善するか

本記事の内容

  • Batch Normalizationの理論と導出
  • Layer Normalizationの理論と導出
  • 学習時と推論時の違い
  • Group Norm, Instance Normとの関係
  • Pythonでの実装と比較実験

前提知識

この記事を読む前に、以下の記事を読んでおくと理解が深まります。

Batch Normalization

基本的なアイデア

Batch Normalization(BatchNorm, BN)は2015年にIoffeとSzegedyによって提案されました。基本的なアイデアは非常にシンプルです。各層の出力をミニバッチ内で平均0・分散1に正規化するのです。

データの前処理で入力を標準化する(平均0・分散1にする)のは常識ですが、BatchNormはこの標準化をネットワークの内部の各層にも適用するアイデアです。

数学的定義

ミニバッチ $\mathcal{B} = \{x_1, x_2, \ldots, x_B\}$ が与えられたとき、BatchNormは以下の4ステップで計算されます。

ステップ1: ミニバッチ平均の計算

$$ \mu_\mathcal{B} = \frac{1}{B}\sum_{i=1}^{B} x_i $$

ステップ2: ミニバッチ分散の計算

$$ \sigma_\mathcal{B}^2 = \frac{1}{B}\sum_{i=1}^{B}(x_i – \mu_\mathcal{B})^2 $$

ステップ3: 正規化

$$ \hat{x}_i = \frac{x_i – \mu_\mathcal{B}}{\sqrt{\sigma_\mathcal{B}^2 + \epsilon}} $$

$\epsilon$ は数値安定性のための小さな定数(通常 $10^{-5}$)です。

ステップ4: スケールとシフト(アフィン変換)

$$ \begin{equation} y_i = \gamma \hat{x}_i + \beta \end{equation} $$

$\gamma$(スケール)と $\beta$(シフト)は学習可能なパラメータです。

なぜスケールとシフトが必要なのか

ステップ3の正規化だけでは、ネットワークの表現力を制限してしまう可能性があります。

たとえば、シグモイド活性化関数の前にBatchNormを置くと、正規化された入力はほぼ $[-2, 2]$ の範囲に収まります。シグモイドはこの範囲でほぼ線形なので、ネットワーク全体がほぼ線形モデルに退化してしまいます。

$\gamma$ と $\beta$ を導入することで、ネットワークは正規化を「元に戻す」ことも含めて、最適な分布を学習できます。極端な場合、$\gamma = \sqrt{\sigma_\mathcal{B}^2 + \epsilon}$, $\beta = \mu_\mathcal{B}$ とすれば、BatchNormは恒等変換になります。つまり、BatchNormはネットワークの表現力を損なわずに、最適化を容易にする仕組みです。

BatchNormの配置

BatchNormは通常、線形変換の後、活性化関数の前に配置されます。

$$ \bm{h} = \sigma(\text{BN}(\bm{W}\bm{x} + \bm{b})) $$

ここでバイアス $\bm{b}$ は不要です。BatchNormの $\beta$ がバイアスの役割を果たすため、$\bm{b}$ を含めても正規化で打ち消されてしまいます。PyTorchでは nn.Linear(in_features, out_features, bias=False) とするのが一般的です。

学習時と推論時の違い

BatchNormは学習時と推論時で振る舞いが異なります。これはBatchNormの実装で最も注意が必要な点です。

学習時: ミニバッチの統計量($\mu_\mathcal{B}$, $\sigma_\mathcal{B}^2$)を使って正規化します。同時に、指数移動平均(Exponential Moving Average, EMA)で全体の統計量を追跡します。

$$ \begin{align} \mu_\text{running} &\leftarrow (1 – m) \cdot \mu_\text{running} + m \cdot \mu_\mathcal{B} \\ \sigma^2_\text{running} &\leftarrow (1 – m) \cdot \sigma^2_\text{running} + m \cdot \sigma^2_\mathcal{B} \end{align} $$

$m$ はモメンタム(通常 $0.1$)です。

推論時: ミニバッチの統計量ではなく、学習中に蓄積した $\mu_\text{running}$ と $\sigma^2_\text{running}$ を使います。これにより、推論時の出力がバッチサイズや他のサンプルに依存しなくなります。

PyTorchでは model.train()model.eval() の切り替えで、この振る舞いが自動的に変わります。model.eval() を忘れると推論結果が不安定になるので注意が必要です。

CNNでのBatchNorm

畳み込みニューラルネットワーク(CNN)では、チャネルごとに正規化を行います。

特徴マップのサイズが $(B, C, H, W)$(バッチ, チャネル, 高さ, 幅)のとき、各チャネル $c$ について $B \times H \times W$ 個の値の平均と分散を計算します。学習パラメータ $\gamma_c$ と $\beta_c$ はチャネルごとに定義されます。

これは「同じフィルタが検出する特徴は、画像内の位置やサンプルによらず同じスケールであるべき」という仮定に基づいています。

BatchNormの仕組みを理解したところで、BatchNormが抱える問題点と、その解決策として提案されたLayer Normalizationを見ていきましょう。

BatchNormの問題点

BatchNormは非常に効果的ですが、いくつかの重要な制限があります。

バッチサイズへの依存: ミニバッチ内の統計量を使うため、バッチサイズが小さいと統計量の推定が不安定になります。バッチサイズ1(オンライン学習)では分散が計算できず、BatchNormは使えません。

逐次データとの相性: RNNやTransformerのように系列データを扱う場合、時刻ごとにバッチ統計量が異なるため、適用が複雑になります。また、推論時に系列長が学習時と異なる場合にも問題が生じます。

分散学習の複雑さ: 複数GPUでの学習では、バッチ統計量をGPU間で同期する必要があり(Synchronized BatchNorm)、通信コストが増加します。

これらの問題を解決するために、バッチ方向ではなく特徴方向で正規化を行うLayer Normalizationが提案されました。

Layer Normalization

基本的なアイデア

Layer Normalization(LayerNorm, LN)は2016年にBaらによって提案されました。BatchNormとの違いは、正規化の方向です。

  • BatchNorm: ミニバッチ内の同じ特徴量について正規化(バッチ方向の統計量)
  • LayerNorm: 1つのサンプル内の全特徴量について正規化(特徴方向の統計量)

イメージとしては、BatchNormは「クラス全体のテストの平均点で正規化する」のに対し、LayerNormは「一人の学生の全科目の平均点で正規化する」ようなものです。

数学的定義

入力 $\bm{x} \in \mathbb{R}^D$($D$ は特徴量の次元数)に対して、LayerNormは以下のように計算されます。

平均と分散の計算(1サンプル内の全特徴量について):

$$ \mu = \frac{1}{D}\sum_{i=1}^{D} x_i, \quad \sigma^2 = \frac{1}{D}\sum_{i=1}^{D}(x_i – \mu)^2 $$

正規化とアフィン変換:

$$ \begin{equation} y_i = \gamma_i \cdot \frac{x_i – \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta_i \end{equation} $$

BatchNormとの重要な違いは以下の点です。

BatchNorm LayerNorm
正規化の方向 バッチ方向($B$ 個のサンプル) 特徴方向($D$ 個の特徴量)
統計量の計算 $\mu, \sigma^2$ はバッチ内で計算 $\mu, \sigma^2$ は各サンプル内で計算
バッチサイズへの依存 あり(小バッチで不安定) なし
学習時と推論時の違い あり(EMA使用) なし(同じ計算)
$\gamma, \beta$ の次元 特徴量ごとに1つ 特徴量ごとに1つ
典型的な用途 CNN Transformer, RNN

LayerNormの利点

バッチサイズに非依存: 1サンプルの内部で正規化するため、バッチサイズが1でも問題なく動作します。

学習と推論で同じ計算: running mean/varianceの管理が不要です。model.train()model.eval() で振る舞いが変わらないため、実装ミスのリスクが減ります。

系列データとの相性: RNNやTransformerでは、各時刻の隠れ状態に対してLayerNormを独立に適用できます。系列長が変わっても問題ありません。

TransformerでのLayerNorm

Transformerアーキテクチャでは、LayerNormが各サブレイヤー(Self-Attention, Feed-Forward)の後に適用されます。

Post-LN(原論文):

$$ \bm{h} = \text{LN}(\bm{x} + \text{SubLayer}(\bm{x})) $$

Pre-LN(GPT-2以降で主流):

$$ \bm{h} = \bm{x} + \text{SubLayer}(\text{LN}(\bm{x})) $$

Pre-LNは学習がより安定し、学習率のウォームアップが不要になることが多いとされています。Pre-LNでは残差接続を通じて勾配が直接流れるパスが保たれるため、深いTransformerでも学習が安定します。

Transformerは系列データを扱い、バッチサイズが比較的小さいことがあるため、BatchNormよりLayerNormが適しています。

では、BatchNormとLayerNormの違いをPythonで実装して、実際にどのように動作するかを確認してみましょう。

Pythonでの実装

BatchNormとLayerNormのスクラッチ実装

import numpy as np
import matplotlib.pyplot as plt

np.random.seed(42)

class BatchNorm1d:
    """Batch Normalization(1次元入力用)"""

    def __init__(self, num_features, momentum=0.1, eps=1e-5):
        self.gamma = np.ones((1, num_features))
        self.beta = np.zeros((1, num_features))
        self.eps = eps
        self.momentum = momentum
        self.running_mean = np.zeros((1, num_features))
        self.running_var = np.ones((1, num_features))
        self.training = True

    def forward(self, X):
        """X: (batch_size, num_features)"""
        if self.training:
            # ミニバッチの統計量
            self.mu = np.mean(X, axis=0, keepdims=True)
            self.var = np.var(X, axis=0, keepdims=True)

            # Running statistics の更新
            self.running_mean = ((1 - self.momentum) * self.running_mean
                                 + self.momentum * self.mu)
            self.running_var = ((1 - self.momentum) * self.running_var
                                + self.momentum * self.var)

            # 正規化
            self.x_hat = (X - self.mu) / np.sqrt(self.var + self.eps)
        else:
            # 推論時はrunning statisticsを使用
            self.x_hat = ((X - self.running_mean)
                          / np.sqrt(self.running_var + self.eps))

        return self.gamma * self.x_hat + self.beta


class LayerNorm:
    """Layer Normalization"""

    def __init__(self, num_features, eps=1e-5):
        self.gamma = np.ones((1, num_features))
        self.beta = np.zeros((1, num_features))
        self.eps = eps

    def forward(self, X):
        """X: (batch_size, num_features)"""
        # 各サンプル内の統計量
        self.mu = np.mean(X, axis=1, keepdims=True)
        self.var = np.var(X, axis=1, keepdims=True)

        # 正規化
        self.x_hat = (X - self.mu) / np.sqrt(self.var + self.eps)

        return self.gamma * self.x_hat + self.beta


# --- デモ: 正規化の効果を可視化 ---
batch_size = 64
num_features = 100

# 偏った分布のデータを生成
X = np.random.randn(batch_size, num_features) * 5 + 3

# 各正規化を適用
bn = BatchNorm1d(num_features)
ln = LayerNorm(num_features)

X_bn = bn.forward(X)
X_ln = ln.forward(X)

# 可視化
fig, axes = plt.subplots(2, 3, figsize=(16, 9))

# 上段: 各特徴量の分布
ax = axes[0, 0]
ax.boxplot([X[:, i] for i in range(0, 100, 10)],
           positions=range(10))
ax.set_title("Original (feature view)", fontsize=12)
ax.set_xlabel("Feature index (sampled)", fontsize=10)
ax.set_ylabel("Value", fontsize=10)
ax.grid(True, alpha=0.3)

ax = axes[0, 1]
ax.boxplot([X_bn[:, i] for i in range(0, 100, 10)],
           positions=range(10))
ax.set_title("After BatchNorm (feature view)", fontsize=12)
ax.set_xlabel("Feature index (sampled)", fontsize=10)
ax.set_ylabel("Value", fontsize=10)
ax.grid(True, alpha=0.3)

ax = axes[0, 2]
ax.boxplot([X_ln[:, i] for i in range(0, 100, 10)],
           positions=range(10))
ax.set_title("After LayerNorm (feature view)", fontsize=12)
ax.set_xlabel("Feature index (sampled)", fontsize=10)
ax.set_ylabel("Value", fontsize=10)
ax.grid(True, alpha=0.3)

# 下段: 各サンプルの分布
ax = axes[1, 0]
ax.boxplot([X[i, :] for i in range(0, 64, 7)],
           positions=range(len(range(0, 64, 7))))
ax.set_title("Original (sample view)", fontsize=12)
ax.set_xlabel("Sample index (sampled)", fontsize=10)
ax.set_ylabel("Value", fontsize=10)
ax.grid(True, alpha=0.3)

ax = axes[1, 1]
ax.boxplot([X_bn[i, :] for i in range(0, 64, 7)],
           positions=range(len(range(0, 64, 7))))
ax.set_title("After BatchNorm (sample view)", fontsize=12)
ax.set_xlabel("Sample index (sampled)", fontsize=10)
ax.set_ylabel("Value", fontsize=10)
ax.grid(True, alpha=0.3)

ax = axes[1, 2]
ax.boxplot([X_ln[i, :] for i in range(0, 64, 7)],
           positions=range(len(range(0, 64, 7))))
ax.set_title("After LayerNorm (sample view)", fontsize=12)
ax.set_xlabel("Sample index (sampled)", fontsize=10)
ax.set_ylabel("Value", fontsize=10)
ax.grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig("bn_vs_ln_distribution.png", dpi=150, bbox_inches="tight")
plt.show()

# 統計量の確認
print("=== 正規化後の統計量 ===")
print(f"BatchNorm - 特徴量ごとの平均: {np.mean(X_bn, axis=0)[:5].round(4)}")
print(f"BatchNorm - 特徴量ごとの分散: {np.var(X_bn, axis=0)[:5].round(4)}")
print(f"LayerNorm - サンプルごとの平均: {np.mean(X_ln, axis=1)[:5].round(4)}")
print(f"LayerNorm - サンプルごとの分散: {np.var(X_ln, axis=1)[:5].round(4)}")

この可視化から、BatchNormとLayerNormの正規化方向の違いが明確にわかります。

  1. 上段(特徴量方向の分布): 元データ(左)では各特徴量の箱ひげ図が同様の分布を示しています(平均3, 分散25程度)。BatchNorm(中央)は各特徴量の分布を平均0・分散1に正規化するため、箱ひげ図が全て同じ位置に揃っています。LayerNorm(右)では特徴量ごとの分布は揃いませんが、サンプル単位で正規化されています

  2. 下段(サンプル方向の分布): BatchNorm(中央)ではサンプルごとの分布にはばらつきが残っています。一方、LayerNorm(右)は各サンプルの全特徴量について平均0・分散1に正規化するため、各サンプルの箱ひげ図が揃っています

  3. 統計量の確認: BatchNormでは特徴量ごとの平均が0、分散が1に近いことが確認でき、LayerNormではサンプルごとの平均が0、分散が1に近いことが確認できます

バッチサイズの影響を検証する

import numpy as np
import matplotlib.pyplot as plt

np.random.seed(42)

# --- バッチサイズがBatchNorm/LayerNormに与える影響 ---
num_features = 50
true_mean = 3.0
true_std = 5.0

batch_sizes = [2, 4, 8, 16, 32, 64, 128, 256]
n_trials = 200

bn_errors = []
ln_errors = []

for bs in batch_sizes:
    bn_var_errors = []
    ln_var_errors = []

    for _ in range(n_trials):
        X = np.random.randn(bs, num_features) * true_std + true_mean

        # BatchNorm: バッチ内の分散
        bn_var = np.var(X, axis=0)  # 各特徴量の分散
        bn_var_error = np.mean(np.abs(bn_var - true_std**2) / true_std**2)

        # LayerNorm: サンプル内の分散
        ln_var = np.var(X, axis=1)  # 各サンプルの分散
        ln_var_error = np.mean(np.abs(ln_var - true_std**2) / true_std**2)

        bn_var_errors.append(bn_var_error)
        ln_var_errors.append(ln_var_error)

    bn_errors.append((np.mean(bn_var_errors), np.std(bn_var_errors)))
    ln_errors.append((np.mean(ln_var_errors), np.std(ln_var_errors)))

fig, axes = plt.subplots(1, 2, figsize=(14, 5.5))

# (a) 分散推定の相対誤差
ax = axes[0]
bn_means = [e[0] for e in bn_errors]
bn_stds = [e[1] for e in bn_errors]
ln_means = [e[0] for e in ln_errors]
ln_stds = [e[1] for e in ln_errors]

ax.errorbar(batch_sizes, bn_means, yerr=bn_stds, fmt="o-", linewidth=2,
            capsize=4, color="tab:blue", label="BatchNorm (batch stat)")
ax.errorbar(batch_sizes, ln_means, yerr=ln_stds, fmt="s-", linewidth=2,
            capsize=4, color="tab:orange", label="LayerNorm (sample stat)")
ax.set_xlabel("Batch size", fontsize=12)
ax.set_ylabel("Relative error of variance", fontsize=12)
ax.set_title("Variance Estimation Error vs Batch Size", fontsize=13)
ax.set_xscale("log", base=2)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# (b) 正規化後の分布の安定性
ax = axes[1]
small_bs_outputs = []
large_bs_outputs = []

for _ in range(100):
    # 小バッチ
    X_small = np.random.randn(4, num_features) * true_std + true_mean
    mu_s = np.mean(X_small, axis=0)
    var_s = np.var(X_small, axis=0)
    X_small_bn = (X_small - mu_s) / np.sqrt(var_s + 1e-5)
    small_bs_outputs.append(np.std(X_small_bn, axis=0))

    # 大バッチ
    X_large = np.random.randn(256, num_features) * true_std + true_mean
    mu_l = np.mean(X_large, axis=0)
    var_l = np.var(X_large, axis=0)
    X_large_bn = (X_large - mu_l) / np.sqrt(var_l + 1e-5)
    large_bs_outputs.append(np.std(X_large_bn, axis=0))

small_stds = np.array(small_bs_outputs).flatten()
large_stds = np.array(large_bs_outputs).flatten()

ax.hist(small_stds, bins=50, density=True, alpha=0.6, color="tab:red",
        label="BN with batch_size=4")
ax.hist(large_stds, bins=50, density=True, alpha=0.6, color="tab:green",
        label="BN with batch_size=256")
ax.axvline(1.0, color="black", linestyle="--", linewidth=2,
           label="Target std = 1.0")
ax.set_xlabel("Std of normalized output", fontsize=12)
ax.set_ylabel("Density", fontsize=12)
ax.set_title("BN Output Stability by Batch Size", fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

plt.tight_layout()
plt.savefig("bn_batch_size_effect.png", dpi=150, bbox_inches="tight")
plt.show()

この実験から、バッチサイズがBatchNormに与える影響が定量的に確認できます。

  1. 左図(分散推定の相対誤差): BatchNorm(青)の誤差はバッチサイズが小さいほど大きく、バッチサイズ2〜4では相対誤差が非常に大きくなっています。バッチサイズが増えるにつれて誤差は減少し、256程度でほぼ安定します。一方、LayerNorm(オレンジ)の誤差はバッチサイズに依存せず一定です。これは50次元の特徴量を使って正規化しているためで、バッチサイズが変わっても影響を受けません

  2. 右図(正規化後の標準偏差の分布): バッチサイズ4(赤)では正規化後の標準偏差が広く分布しており、目標の1.0から大きくずれることがあります。バッチサイズ256(緑)ではほぼ1.0に集中しており、正規化が安定しています。これはバッチサイズが小さいとミニバッチの統計量がサンプルノイズの影響を受けやすいためです

他の正規化手法

Group Normalization

Group Normalization(GroupNorm, GN)は2018年にWuとHeによって提案されました。LayerNormとBatchNormの中間に位置する手法です。

チャネルを $G$ 個のグループに分割し、グループ内で正規化を行います。

$$ \mu_g = \frac{1}{|S_g|}\sum_{(c,h,w) \in S_g} x_{c,h,w}, \quad \sigma_g^2 = \frac{1}{|S_g|}\sum_{(c,h,w) \in S_g} (x_{c,h,w} – \mu_g)^2 $$

ここで $S_g$ はグループ $g$ に属するチャネルとその空間位置のインデックスです。

  • $G = 1$: LayerNormと同等
  • $G = C$(チャネル数): Instance Normと同等

GroupNormはバッチサイズに依存せず、かつチャネル間の相関を考慮できるため、物体検出やセマンティックセグメンテーションなどバッチサイズが小さくなりがちなタスクで効果的です。

Instance Normalization

Instance Normalization(InstanceNorm, IN)は、各チャネルを個別に正規化します。スタイル変換(Style Transfer)で提案され、画像生成タスクで広く使われています。

RMSNorm

RMS Normalization(RMSNorm)は、LayerNormから平均の引き算を省略し、RMS(Root Mean Square)のみで正規化する手法です。

$$ \begin{equation} \text{RMS}(\bm{x}) = \sqrt{\frac{1}{D}\sum_{i=1}^{D} x_i^2}, \quad y_i = \frac{\gamma_i \cdot x_i}{\text{RMS}(\bm{x})} \end{equation} $$

LLaMAなど最近の大規模言語モデルで採用されており、LayerNormとほぼ同等の性能で計算コストが低いことが報告されています。

正規化手法の使い分けまとめ

手法 正規化方向 バッチ依存 主な用途
BatchNorm バッチ方向 あり CNN(画像分類)
LayerNorm 特徴方向 なし Transformer, RNN
GroupNorm グループ内の特徴方向 なし CNN(小バッチ)
InstanceNorm チャネル内の空間方向 なし スタイル変換
RMSNorm 特徴方向(平均なし) なし LLM(LLaMA等)

正規化がなぜ効くのか — 理論的な解釈

損失曲面の平滑化

Santurkarら(2018)は、BatchNormの効果は内部共変量シフトの解消ではなく、損失曲面(loss landscape)の平滑化にあると主張しました。

BatchNormを使うと、損失関数のリプシッツ定数(勾配の変化の上界)が小さくなり、損失曲面が滑らかになります。これにより、勾配降下法がより大きな学習率で安定して収束でき、学習が加速されます。

勾配の分散の安定化

正規化によって各層の入力分布が安定するため、勾配の分散も安定します。これは特に深いネットワークで重要で、勾配消失・勾配爆発のリスクを低減します。

直感的には、正規化は「損失曲面の等高線を円形に近づける」効果があり、勾配降下法のジグザグ問題(楕円形の等高線でのゆっくりとした収束)を緩和します。

まとめ

本記事では、Batch NormalizationとLayer Normalizationの理論と実装を解説しました。

  • Batch Normalization: ミニバッチ内の統計量で正規化。CNNで標準的に使用される。学習時と推論時で異なる挙動をする点に注意
  • Layer Normalization: 各サンプル内の統計量で正規化。バッチサイズに依存せず、TransformerやRNNで使用される。学習時と推論時で同じ計算
  • BatchNormはバッチサイズが小さいと不安定になるため、小バッチ環境ではLayerNormやGroupNormが適切
  • 正規化層はスケール $\gamma$ とシフト $\beta$ の学習パラメータを持ち、必要に応じて正規化を「元に戻す」ことができる
  • 最近のLLMではRMSNormが採用され、LayerNormから平均の計算を省略して計算効率を改善

次のステップとして、以下の記事も参考にしてください。