知識蒸留の理論と実装 — 大きなモデルの知識を小さなモデルに転写する

1億パラメータのモデルが99%の精度を達成しました。しかし、このモデルをスマートフォンやIoTデバイスに載せるには大きすぎます。推論に必要なメモリは数百MB、レイテンシは数秒に及び、リアルタイムアプリケーションには使えません。パラメータを100分の1にした小さなモデルを一から学習すると、精度は90%まで落ちてしまいます。

「大きなモデルの精度を保ちながら、小さなモデルの効率を実現する」方法はないでしょうか。これは「ベテランの職人の技術を新人に効率的に伝授する」問題に似ています。新人がゼロから全てを学ぶよりも、ベテランの「コツ」や「判断基準」を直接教わった方が、遥かに速く成長できます。知識蒸留(Knowledge Distillation; Hinton et al., 2015)は、機械学習の文脈でまさにこのアプローチを実現する手法です。

知識蒸留の核心的なアイデアは、大きなモデル(教師モデル)のソフト出力(確率分布)を、小さなモデル(生徒モデル)の学習に利用することです。ソフト出力には「正解は猫だが、虎にも似ている」といったクラス間の類似性情報(暗黙知)が含まれており、これがハードラベル(正解のみ1, 他は0)にはない豊富な学習信号を提供します。

たとえば、犬の画像を分類する場合を考えましょう。ハードラベルは単に「犬」と教えるだけですが、教師モデルのソフト出力は「犬82%、オオカミ8%、猫5%、馬3%、…」という情報を含みます。この出力から生徒モデルは「犬とオオカミは視覚的に似ている」「犬と馬はあまり似ていない」といったクラス間の構造を学ぶことができます。1枚の画像から得られる情報量が、ハードラベルに比べて格段に増えるのです。

知識蒸留を理解すると、以下のような場面で活用できます。

  • モデル圧縮: エッジデバイスへの高精度モデルのデプロイ。スマートフォン、ドローン、IoTセンサーなど、計算資源が限られた環境での推論を可能にします
  • 推論の高速化: レイテンシが重要なリアルタイムアプリケーション。自動運転の物体検出や、チャットボットの応答生成など
  • アンサンブルの圧縮: 複数モデルのアンサンブルを1つのモデルに蒸留。アンサンブルは精度が高いものの推論コストが高いため、蒸留で1つのモデルに知識を集約できます
  • データプライバシー: 元のデータにアクセスせずにモデルの知識を転写。教師モデルの出力さえあれば蒸留が可能なため、データの共有が制限される場面で活用できます

本記事の内容

  • 知識蒸留の動機とソフトターゲットの直感
  • 温度付きソフトマックスの理論
  • 蒸留損失関数の数学的定式化
  • 温度パラメータの影響
  • Pythonでの実装と実験

前提知識

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

ハードラベル vs ソフトラベル

ハードラベルの限界

通常の教師あり学習では、正解ラベルをワンホットベクトル(ハードラベル)で表現します。

犬の画像のラベル: $\bm{y} = [0, 0, 1, 0, 0, \ldots]$(犬クラスのみ1)

このハードラベルは「この画像は犬である」という事実しか伝えません。情報理論的に見ると、$K$ クラス分類のハードラベルが持つ情報量は $\log_2 K$ ビットです。10クラス分類なら約3.3ビットにすぎません。しかし実際には

  • この犬はオオカミに似ているか?
  • 猫との区別は容易か?
  • 馬との類似性はどの程度か?

といったクラス間の関係性の情報は失われています。一方、ソフトラベルは $K$ 個の連続値を含むため、遥かに多くの情報を持ちます。

ソフトラベルの情報量

一方、訓練済みの大きなモデルの出力確率

$$ \bm{p} = [0.01, 0.05, 0.82, 0.02, 0.08, \ldots] $$

は、「犬が82%で最も確率が高いが、オオカミにも8%の確率がある」という豊かな情報を含みます。

Hinton et al.(2015)はこの確率分布をダークナレッジ(dark knowledge)と呼びました。正解ラベルには現れないが、モデルが学習した「クラス間の構造」がソフト出力に暗黙的に符号化されています。

小さなモデルがこのソフト出力から学習すると、ハードラベルだけで学習するよりも多くの情報を得られ、より高い精度を達成できます。各訓練サンプルから得られる学習信号がリッチになるため、少ないデータでも効率的な学習が可能です。

しかし、ここで1つ問題があります。訓練済みの大きなモデルの通常のソフトマックス出力では、正解クラスの確率が非常に高く(0.99等)、他のクラスの確率はほぼ0です。これではハードラベルとほとんど変わりません。この極端な分布からは有用な情報を引き出しにくいため、温度パラメータで分布を「ソフト」にする工夫が必要です。次に、この温度パラメータの仕組みを詳しく見ていきましょう。

温度付きソフトマックス

定義

通常のソフトマックスは、ロジット $z_i$(最終全結合層の出力、つまりソフトマックスの前の値)に対して

$$ p_i = \frac{\exp(z_i)}{\sum_j \exp(z_j)} $$

です。この標準的なソフトマックスは $T = 1$ の場合に相当します。知識蒸留では温度 $T > 0$ を導入して、分布の「鋭さ」を制御します。

$$ \begin{equation} q_i = \frac{\exp(z_i / T)}{\sum_j \exp(z_j / T)} \end{equation} $$

温度の直感

温度 $T$ はソフトマックス出力の「鋭さ」を制御します。

  • $T \to 0$: 最大ロジットのクラスに確率が集中($\arg\max$ に近い)
  • $T = 1$: 通常のソフトマックス
  • $T \to \infty$: 一様分布に近づく(全クラスがほぼ等確率)

この直感を定量的に理解しましょう。ロジットが $z_1 = 5, z_2 = 3, z_3 = 1$ のとき

$T$ $q_1$ $q_2$ $q_3$
0.5 0.998 0.002 0.000
1.0 0.844 0.114 0.042
2.0 0.576 0.272 0.152
5.0 0.419 0.330 0.251

$T = 1$ では $q_1$ が支配的ですが、$T = 5$ ではクラス間の差が緩やかになり、「$z_2$ が $z_3$ より大きい」という相対的な関係がより明確に読み取れます。

「温度」という名称は統計力学に由来します。統計力学のボルツマン分布 $p_i \propto \exp(-E_i / k_B T)$ では、温度が高いほどエネルギー準位間の差が小さくなり、全ての状態が等確率に近づきます。知識蒸留の温度パラメータも同じ原理で、$T$ を上げるとロジット間の差が相対的に縮小し、均一な分布に近づきます。

なぜ温度を上げるのか

温度を上げる理由は、小さい確率値に含まれる情報を引き出すためです。

$T = 1$ のとき、正解クラスの確率が0.99で他が0.005程度だと、KLダイバージェンスの勾配はほぼ正解クラスにしか伝わりません。$T = 5$ にすると、例えば0.30と0.20の差が勾配として伝わり、「このクラスの方が正解に近い」という情報が生徒モデルに届きます。

温度と勾配の関係

Hinton et al.(2015)は、$T$ が十分大きいとき、蒸留損失のKLダイバージェンス項の勾配が

$$ \frac{\partial}{\partial z_i^S}\text{KL}(q^T \| q^S) \approx \frac{1}{T^2}\left(\frac{z_i^S}{T} – \frac{z_i^T}{T}\right) $$

と近似できることを示しました。ここで上付きの $T$ と $S$ はそれぞれ教師(Teacher)と生徒(Student)を表します。

この近似から、蒸留損失は教師と生徒のロジットの差 $z_i^S – z_i^T$ に比例する勾配を与えることがわかります。つまり、生徒のロジットが教師のロジットに近づくように学習されるのです。

$T^2$ で割られるため、温度が高いほど勾配が小さくなります。このスケーリングを補正するため、蒸留損失に $T^2$ を掛けるのが標準的な実装です。

この近似をもう少し詳しく見てみましょう。$T$ が十分大きいとき、$z_i / T$ は小さくなるため、$\exp(z_i / T) \approx 1 + z_i / T$ と近似できます。これを正規化すると

$$ q_i \approx \frac{1 + z_i / T}{\sum_j (1 + z_j / T)} = \frac{1}{K} + \frac{z_i}{KT} + O(T^{-2}) $$

ここで $K$ はクラス数です。つまり、高温では各クラスの確率は一様分布 $1/K$ からのロジットに比例した微小なずれで表されます。この近似の下で KL ダイバージェンスの勾配を計算すると、ロジット間の差に比例する形が自然に導かれます。

温度付きソフトマックスの仕組みが明らかになったところで、次に蒸留の損失関数全体の定式化を見ていきましょう。

蒸留損失関数

定式化

知識蒸留の損失関数は、2つの項の加重和です。

$$ \begin{equation} \mathcal{L} = (1 – \alpha) \cdot \mathcal{L}_{\text{hard}} + \alpha \cdot T^2 \cdot \mathcal{L}_{\text{soft}} \end{equation} $$

ハードラベル損失: 生徒モデルの出力($T=1$)と正解ラベルの交差エントロピー

$$ \mathcal{L}_{\text{hard}} = -\sum_i y_i \log p_i^S $$

ソフトラベル損失: 生徒モデルの出力(温度 $T$)と教師モデルの出力(温度 $T$)のKLダイバージェンス

$$ \mathcal{L}_{\text{soft}} = \text{KL}(q^T \| q^S) = \sum_i q_i^T \log \frac{q_i^T}{q_i^S} $$

ここで $q_i^T, q_i^S$ はそれぞれ教師・生徒の温度 $T$ でのソフトマックス出力です。

各項の役割

  • $\mathcal{L}_{\text{hard}}$: 正解ラベルとの整合性を保証。教師モデルが完璧ではない場合に、生徒が正解ラベルからも直接学習することで最低限の精度を達成するための「安全装置」です
  • $\mathcal{L}_{\text{soft}}$: 教師の知識(クラス間の関係性)を転写。温度 $T$ で「柔らかくした」分布から学習します。この項が知識蒸留の核心であり、ハードラベルにはないクラス間の構造情報を生徒に伝えます
  • $T^2$: 温度による勾配のスケーリングを補正。先ほど見たように、KLダイバージェンスの勾配は $1/T^2$ に比例して小さくなるため、$T^2$ を掛けることで適切な勾配のスケールを維持します。この補正がないと、$T$ を上げたときに蒸留損失の勾配が無視できるほど小さくなってしまいます

$\alpha$ のバランス

$\alpha \in [0, 1]$ はハードラベルとソフトラベルのバランスを制御します。

  • $\alpha = 0$: 通常の教師あり学習(蒸留なし)
  • $\alpha = 1$: 教師のソフト出力のみから学習
  • 典型的な設定: $\alpha = 0.5$ 〜 $0.9$

Hinton et al.の実験では、$\alpha$ を大きめ(教師の知識を重視)にした方が良い結果が得られることが多いと報告されています。ただし、教師モデルの精度が低い場合は $\alpha$ を小さくする必要があります。なぜなら、精度の低い教師のソフト出力は「誤った」クラス間関係を含んでおり、それを重視すると生徒の性能が低下するためです。

$\alpha$ の選択に迷った場合は、まず $\alpha = 0.5$ で実験を始め、教師の精度が十分に高い場合は $\alpha$ を徐々に大きくしていくアプローチが実践的です。

温度 $T$ の選択

  • $T = 1$: 蒸留の効果が小さい
  • $T = 3$-$5$: 一般的に良好な結果
  • $T = 10$-$20$: 非常にソフトな分布。教師の精度が高い場合に有効
  • $T$ が大きすぎる: 一様分布に近づき、情報が失われる

実用的には $T = 3$-$5$ が出発点として推奨されます。

$T$ と $\alpha$ は互いに関連しています。$T$ を大きくするとソフトラベルの情報がより均一になるため、$\alpha$ を大きくして蒸留損失の重みを増やす方がよい結果が得られる傾向があります。逆に、$T$ が小さい場合はソフトラベルがハードラベルに近いため、$\alpha$ を大きくする必要性は低くなります。

理論を理解したところで、Pythonで実装してみましょう。教師モデルの学習、ハードラベルのみでの生徒モデルの学習、そして知識蒸留を用いた生徒モデルの学習を比較します。

Pythonでの実装

教師-生徒モデルの蒸留

NumPyで簡単な2層ニューラルネットワークを使った蒸留実験を行います。教師モデルは隠れ層128ユニット、生徒モデルは隠れ層32ユニットとし、パラメータ数に約4倍の差がある状況で蒸留の効果を確認します。

import numpy as np
import matplotlib.pyplot as plt

def softmax(z, T=1.0):
    """温度付きソフトマックス"""
    z_scaled = z / T
    e = np.exp(z_scaled - z_scaled.max(axis=-1, keepdims=True))
    return e / e.sum(axis=-1, keepdims=True)

def cross_entropy(p, q, eps=1e-10):
    """交差エントロピー"""
    return -np.sum(p * np.log(q + eps), axis=-1).mean()

def kl_divergence(p, q, eps=1e-10):
    """KLダイバージェンス"""
    return np.sum(p * np.log((p + eps) / (q + eps)), axis=-1).mean()

class SimpleNN:
    """シンプルな2層ニューラルネットワーク"""
    def __init__(self, input_dim, hidden_dim, output_dim):
        scale = np.sqrt(2.0 / input_dim)
        self.W1 = np.random.randn(input_dim, hidden_dim) * scale
        self.b1 = np.zeros(hidden_dim)
        self.W2 = np.random.randn(hidden_dim, output_dim) * np.sqrt(2.0 / hidden_dim)
        self.b2 = np.zeros(output_dim)

    def forward(self, X):
        """順伝播"""
        self.h1 = np.maximum(0, X @ self.W1 + self.b1)  # ReLU
        self.logits = self.h1 @ self.W2 + self.b2
        return self.logits

    def predict(self, X):
        logits = self.forward(X)
        return softmax(logits)

    def backward(self, X, grad_logits, lr=0.01):
        """逆伝播"""
        batch_size = X.shape[0]
        dW2 = self.h1.T @ grad_logits / batch_size
        db2 = grad_logits.mean(axis=0)
        dh1 = grad_logits @ self.W2.T
        dh1 = dh1 * (self.h1 > 0)  # ReLU grad
        dW1 = X.T @ dh1 / batch_size
        db1 = dh1.mean(axis=0)

        np.clip(dW2, -1, 1, out=dW2)
        np.clip(dW1, -1, 1, out=dW1)
        self.W2 -= lr * dW2
        self.b2 -= lr * db2
        self.W1 -= lr * dW1
        self.b1 -= lr * db1
np.random.seed(42)

# 合成データ: 10クラス分類
n_samples = 2000
input_dim = 20
n_classes = 10
hidden_teacher = 128
hidden_student = 32

# データ生成
W_true = np.random.randn(input_dim, n_classes)
X = np.random.randn(n_samples, input_dim)
logits_true = X @ W_true
y = np.argmax(logits_true, axis=1)
Y_onehot = np.eye(n_classes)[y]

# 訓練・テスト分割
split = int(0.8 * n_samples)
X_train, X_test = X[:split], X[split:]
Y_train, Y_test = Y_onehot[:split], Y_onehot[split:]
y_train, y_test = y[:split], y[split:]

# 教師モデル(大きい)の学習
teacher = SimpleNN(input_dim, hidden_teacher, n_classes)
teacher_losses = []
for epoch in range(100):
    perm = np.random.permutation(len(X_train))
    epoch_loss = 0
    for start in range(0, len(X_train) - 64, 64):
        idx = perm[start:start+64]
        logits = teacher.forward(X_train[idx])
        probs = softmax(logits)
        loss = cross_entropy(Y_train[idx], probs)
        grad = (probs - Y_train[idx]) / 64
        teacher.backward(X_train[idx], grad, lr=0.05)
        epoch_loss += loss
    teacher_losses.append(epoch_loss / (len(X_train) // 64))

teacher_acc = (np.argmax(teacher.predict(X_test), axis=1) == y_test).mean()
print(f"Teacher accuracy: {teacher_acc:.1%}")

# 生徒モデル1: ハードラベルのみで学習
student_hard = SimpleNN(input_dim, hidden_student, n_classes)
hard_losses = []
for epoch in range(100):
    perm = np.random.permutation(len(X_train))
    epoch_loss = 0
    for start in range(0, len(X_train) - 64, 64):
        idx = perm[start:start+64]
        logits = student_hard.forward(X_train[idx])
        probs = softmax(logits)
        loss = cross_entropy(Y_train[idx], probs)
        grad = (probs - Y_train[idx]) / 64
        student_hard.backward(X_train[idx], grad, lr=0.05)
        epoch_loss += loss
    hard_losses.append(epoch_loss / (len(X_train) // 64))

hard_acc = (np.argmax(student_hard.predict(X_test), axis=1) == y_test).mean()
print(f"Student (hard labels): {hard_acc:.1%}")

# 生徒モデル2: 知識蒸留で学習
T = 4.0  # 温度
alpha = 0.7  # ソフトラベルの重み

student_kd = SimpleNN(input_dim, hidden_student, n_classes)
kd_losses = []

for epoch in range(100):
    perm = np.random.permutation(len(X_train))
    epoch_loss = 0
    for start in range(0, len(X_train) - 64, 64):
        idx = perm[start:start+64]

        # 教師のソフトターゲット
        teacher_logits = teacher.forward(X_train[idx])
        teacher_soft = softmax(teacher_logits, T=T)

        # 生徒の出力
        student_logits = student_kd.forward(X_train[idx])
        student_soft = softmax(student_logits, T=T)
        student_hard_prob = softmax(student_logits, T=1.0)

        # 蒸留損失の勾配
        grad_soft = (student_soft - teacher_soft) / T  # KLの勾配(T^2は後で掛ける)
        grad_hard = (student_hard_prob - Y_train[idx])

        grad = (alpha * T**2 * grad_soft + (1 - alpha) * grad_hard) / 64
        student_kd.backward(X_train[idx], grad, lr=0.05)

        loss_soft = kl_divergence(teacher_soft, student_soft)
        loss_hard = cross_entropy(Y_train[idx], student_hard_prob)
        epoch_loss += alpha * T**2 * loss_soft + (1 - alpha) * loss_hard
    kd_losses.append(epoch_loss / (len(X_train) // 64))

kd_acc = (np.argmax(student_kd.predict(X_test), axis=1) == y_test).mean()
print(f"Student (distillation): {kd_acc:.1%}")
fig, axes = plt.subplots(1, 3, figsize=(18, 5))

# (a) 学習曲線
ax = axes[0]
ax.plot(teacher_losses, "b-", linewidth=1.5, label=f"Teacher (h={hidden_teacher})")
ax.plot(hard_losses, "r--", linewidth=1.5, label=f"Student-Hard (h={hidden_student})")
ax.plot(kd_losses, "g-.", linewidth=1.5, label=f"Student-KD (h={hidden_student})")
ax.set_xlabel("Epoch", fontsize=12)
ax.set_ylabel("Loss", fontsize=12)
ax.set_title("Training Loss", fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# (b) 温度による分布の変化
ax = axes[1]
sample_logits = teacher.forward(X_test[:1])[0]
temps = [0.5, 1.0, 3.0, 5.0, 10.0]
colors = plt.cm.viridis(np.linspace(0, 1, len(temps)))
x_pos = np.arange(n_classes)
for temp, color in zip(temps, colors):
    probs = softmax(sample_logits.reshape(1, -1), T=temp)[0]
    ax.plot(x_pos, probs, "o-", color=color, linewidth=1.5,
            markersize=5, label=f"T={temp}")
ax.set_xlabel("Class", fontsize=12)
ax.set_ylabel("Probability", fontsize=12)
ax.set_title("Temperature Effect on Softmax", fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)

# (c) 精度比較
ax = axes[2]
models = ["Teacher\n(large)", "Student\n(hard)", "Student\n(distilled)"]
accs = [teacher_acc, hard_acc, kd_acc]
colors = ["#8da0cb", "#fc8d62", "#66c2a5"]
bars = ax.bar(models, [a * 100 for a in accs], color=colors,
              edgecolor="black", linewidth=0.8)
for bar, acc in zip(bars, accs):
    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.5,
            f"{acc:.1%}", ha='center', fontsize=11, fontweight='bold')
ax.set_ylabel("Accuracy (%)", fontsize=12)
ax.set_title(f"Accuracy Comparison (T={T}, $\\alpha$={alpha})", fontsize=13)
ax.grid(True, alpha=0.3, axis="y")
ax.set_ylim(0, 100)

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

このグラフから、知識蒸留の効果が確認できます。

  1. 左図(学習曲線): 教師モデル(青)が最も低い損失に到達し、ハードラベルの生徒(赤点線)と蒸留された生徒(緑一点鎖線)が続きます。蒸留された生徒はハードラベルの生徒よりも速く収束する傾向があり、教師のソフトターゲットが効率的な学習信号を提供していることを示しています

  2. 中央図(温度の効果): 同じロジットに対して温度を変えたソフトマックス出力です。$T = 0.5$ では正解クラスにほぼ全確率が集中し、$T = 10$ ではほぼ一様分布になります。$T = 3$-$5$ 付近が、クラス間の相対的な関係を保ちつつ十分に「ソフト」な分布を提供しており、蒸留に適していることがわかります

  3. 右図(精度比較): 教師モデルが最も高い精度を達成し、蒸留された生徒がハードラベルの生徒を上回っています。重要なのは、蒸留された生徒は教師の4分の1のパラメータ数で、ハードラベルの生徒より高い精度を達成していることです。これが知識蒸留の実用的な価値です

この実験では合成データを使用しているため、差は控えめですが、実際の画像分類や自然言語処理のタスクでは、蒸留による精度向上はさらに顕著になることが報告されています。たとえば、BERT-Base(110Mパラメータ)の知識をDistilBERT(66Mパラメータ)に蒸留した場合、パラメータ数を40%削減しながらも、元のBERTの97%の精度を維持できたことが知られています。

蒸留がなぜ機能するか — 理論的な解釈

蒸留が機能する理由をもう少し深く考えてみましょう。ハードラベルでの学習では、各サンプルから得られる勾配情報は基本的に「正解クラスの確率を上げ、他のクラスの確率を下げる」という方向だけです。しかし、ソフトラベルでの学習では、「クラス2の確率をクラス3より高くする」「クラス5とクラス7の確率はほぼ同じにする」など、クラス間の相対的な関係に関する勾配情報も得られます。

これは教育のアナロジーで言えば、「答えは A です」と教えるのと、「A が正解ですが、B も少し正しい要素があり、C は全く間違いです」と教えるのの違いです。後者の方が学習者は概念の構造をより深く理解できます。

数学的には、ソフトラベルによる学習はラベルスムージングの一般化と見なすこともできます。ラベルスムージングは全クラスに均等に確率を分配しますが、蒸留のソフトラベルはデータに依存した「賢い」確率分配を行うため、より有用な正則化効果を持ちます。

蒸留の基本的な仕組みを理解したところで、この手法の様々な変種と発展を見ていきましょう。

蒸留の変種と発展

知識蒸留の基本形(Hinton et al., 2015)に対して、多くの発展的な手法が提案されています。ここでは代表的な3つの変種を紹介します。

Feature-based Distillation

Romero et al.(2015)のFitNetsは、最終出力だけでなく中間層の特徴量を蒸留します。

$$ \mathcal{L}_{\text{feature}} = \|\bm{h}^S – r(\bm{h}^T)\|^2 $$

ここで $\bm{h}^T, \bm{h}^S$ は教師・生徒の中間層の出力、$r$ は次元を合わせるための線形変換(教師と生徒の隠れ層の次元が異なる場合に必要)です。中間表現を直接近づけることで、より深い知識の転写が可能です。

最終出力のソフトラベルが「何を予測するか」の知識を伝えるのに対し、中間層の特徴量は「どのような表現を学ぶか」の知識を伝えます。たとえば、画像分類の場合、中間層の特徴量にはエッジ検出やテクスチャ認識といった低レベルの視覚特徴が含まれており、これらを直接転写することで生徒モデルの表現学習を効率化できます。

Self-Distillation

教師モデルなしで、モデル自身の過去の出力を蒸留に使う手法です。学習の初期段階の予測をソフトターゲットとして使うことで、モデルの正則化効果が得られます。Born-Again Networks(Furlanello et al., 2018)では、同じアーキテクチャの「世代」を重ねることで精度が向上します。

具体的には、モデルAを通常通り学習した後、モデルAのソフト出力を教師としてモデルB(同じアーキテクチャ)を蒸留で学習します。驚くべきことに、モデルBはモデルAよりも高い精度を達成することがあります。これは、ソフトラベルが「データの難しさ」に関する情報を含んでおり、正則化として機能するためと考えられています。

Online Distillation

教師と生徒を同時に学習する手法です。Deep Mutual Learning(Zhang et al., 2018)では、2つのモデルが互いの出力をソフトターゲットとして学習し合います。事前に訓練済みの教師が不要なため、実用的です。

Online Distillation の損失関数は各モデルについて以下のように定式化されます。

$$ \mathcal{L}_1 = \mathcal{L}_{\text{CE}}(\bm{p}_1, \bm{y}) + \text{KL}(\bm{p}_2 \| \bm{p}_1) $$

$$ \mathcal{L}_2 = \mathcal{L}_{\text{CE}}(\bm{p}_2, \bm{y}) + \text{KL}(\bm{p}_1 \| \bm{p}_2) $$

2つのモデルが互いに「教え合う」ことで、どちらも単独で学習するよりも高い精度に到達します。この手法は、大きな教師モデルを事前に学習する計算コストを省けるため、計算資源が限られた場面で特に有効です。

蒸留手法の比較

手法 教師の要件 転写される知識 主な利点
Response-based(標準的蒸留) 訓練済み大規模モデル ソフト出力(クラス間関係) 実装が簡単、理論的裏付け
Feature-based(FitNets) 訓練済み大規模モデル 中間層の特徴量 より深い知識転写
Self-Distillation 不要(自身) 過去の自分のソフト出力 教師不要、正則化効果
Online Distillation 不要(相互) 互いのソフト出力 教師不要、両方が改善

まとめ

本記事では、知識蒸留の理論と実装について解説しました。

  • 知識蒸留は教師モデルのソフト出力(ダークナレッジ)を生徒モデルの学習に利用し、モデル圧縮と精度の両立を実現する。ソフト出力にはクラス間の類似性情報が含まれており、ハードラベルよりもリッチな学習信号を提供する
  • 温度パラメータ $T$ はソフトマックス出力の鋭さを制御し、$T = 3$-$5$ が一般的な設定。高温ではロジット間の差が緩やかになり、クラス間の相対的な関係がより明確に伝わる
  • 蒸留損失は $\mathcal{L} = (1-\alpha)\mathcal{L}_{\text{hard}} + \alpha T^2 \mathcal{L}_{\text{soft}}$ で、$T^2$ が勾配のスケーリングを補正する。$\alpha$ は教師の精度に応じて調整する
  • 蒸留された生徒は同じアーキテクチャのハードラベル学習よりも高い精度を達成できる。DistilBERTのように、パラメータ数を40%削減しながら元のモデルの97%の精度を維持した実例もある
  • Feature-based、Self-Distillation、Online Distillationなどの発展的な手法により、蒸留の適用範囲はさらに広がっている

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