Test-Time Augmentation × EMAの組み合わせで精度はさらに伸びるか?【Keras×CIFAR-10実験】

投稿日:2026年8月8日土曜日 最終更新日:

CIFAR-10 CNN EMA Google Colab Keras TTA 過学習 画像分類

X f B! L
Test-Time Augmentation × EMAの組み合わせで精度はさらに伸びるか?【Keras×CIFAR-10実験】 アイキャッチ画像

EMA(重みの指数移動平均)とTTA(Test-Time Augmentation)は、それぞれ単体の効果は以前の記事で確認済みです。今回はこの2つを組み合わせたら精度はさらに伸びるのか、それとも片方だけで十分なのかをCIFAR-10で検証します。

EMAは学習中の重みを指数移動平均で滑らかにする「重み空間」でのアンサンブル的手法、TTAは推論時に複数の水増し画像を平均する「予測空間」でのアンサンブル的手法です。作用する場所が違うため、加算的に効く(両方使うと相乗効果)のか、片方が既にカバーしている分散をもう片方が重複してカバーするだけ(効果が頭打ち)なのかは、実際に試してみないとわかりません。

📘 この記事でわかること

  • EMA単体・TTA単体・組み合わせでtest_accuracyはどう変わるか
  • 「重み空間の平均化」と「予測空間の平均化」は加算的に効くのか、それとも頭打ちになるのか
  • KerasのEMA自動終了処理・TTA実装時にハマりやすいポイント

EMAとTTAはそれぞれ何を平均化しているのか

EMAは、各ステップの重みθを指数移動平均でなだらかにします。

$$ \theta_{EMA} \leftarrow \beta \, \theta_{EMA} + (1-\beta)\, \theta $$

これにより、学習終盤の重みの揺れ(特定のミニバッチへの過適合)が平均化され、より汎化性能の高い重みが得られます。Kerasではoptimizeruse_ema=Trueを指定するだけで有効化でき、model.fit()終了時に自動的にEMA重みへ切り替わります(finalize_variable_values()を手動で呼ぶ必要はありません)。

一方TTAは、推論時に入力画像に複数の変換T_k(水平反転など)を施し、それぞれの予測確率を平均します。

$$ \hat{y} = \frac{1}{K}\sum_{k=1}^{K} f_\theta\bigl(T_k(x)\bigr) $$

EMAは「学習中に得られた複数の重み」を平均し、TTAは「1つの重みに対する複数の入力」を平均します。平均化の対象(重み vs 入力)が異なるため、理屈の上では独立に効果を持ちやすいはずですが、両方とも最終的には「モデルの予測のブレを減らす」という同じ目的地に向かっているため、重複が生じる可能性もあります。

パターンEMATTA平均化の対象
A:baselineなしなしなし
B:EMAのみありなし重み空間
C:TTAのみなしあり予測空間
D:EMA+TTAありあり重み空間+予測空間

実験コード

環境準備(最初に一度だけ実行)

# ── 環境準備(最初に一度だけ実行)──────────────────────
!apt-get -y install fonts-ipafont-gothic
!rm -rf /root/.cache/matplotlib
!pip install -q japanize_matplotlib
print("環境準備完了")
実行結果をクリックして内容を開く
Reading package lists... Done
Building dependency tree... Done
Reading state information... Done
The following additional packages will be installed:
  fonts-ipafont-mincho
The following NEW packages will be installed:
  fonts-ipafont-gothic fonts-ipafont-mincho
0 upgraded, 2 newly installed, 0 to remove and 53 not upgraded.
Need to get 8,237 kB of archives.
After this operation, 28.7 MB of additional disk space will be used.
Get:1 http://archive.ubuntu.com/ubuntu jammy/universe amd64 fonts-ipafont-gothic all 00303-21ubuntu1 [3,513 kB]
Get:2 http://archive.ubuntu.com/ubuntu jammy/universe amd64 fonts-ipafont-mincho all 00303-21ubuntu1 [4,724 kB]
Fetched 8,237 kB in 2s (4,823 kB/s)
Selecting previously unselected package fonts-ipafont-gothic.
(Reading database ... 122403 files and directories currently installed.)
Preparing to unpack .../fonts-ipafont-gothic_00303-21ubuntu1_all.deb ...
Unpacking fonts-ipafont-gothic (00303-21ubuntu1) ...
Selecting previously unselected package fonts-ipafont-mincho.
Preparing to unpack .../fonts-ipafont-mincho_00303-21ubuntu1_all.deb ...
Unpacking fonts-ipafont-mincho (00303-21ubuntu1) ...
Setting up fonts-ipafont-mincho (00303-21ubuntu1) ...
update-alternatives: using /usr/share/fonts/opentype/ipafont-mincho/ipam.ttf to provide /usr/share/fonts/truetype/fonts-japanese-mincho.ttf (fonts-japanese-mincho.ttf) in auto mode
Setting up fonts-ipafont-gothic (00303-21ubuntu1) ...
update-alternatives: using /usr/share/fonts/opentype/ipafont-gothic/ipag.ttf to provide /usr/share/fonts/truetype/fonts-japanese-gothic.ttf (fonts-japanese-gothic.ttf) in auto mode
Processing triggers for fontconfig (2.13.1-4.2ubuntu5) ...
     ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 4.1/4.1 MB 54.2 MB/s eta 0:00:00
  Preparing metadata (setup.py) ... done
  Building wheel for japanize_matplotlib (setup.py) ... done
環境準備完了

import・データ準備・モデル構築関数

import tensorflow as tf
from tensorflow import keras
import matplotlib.pyplot as plt
import japanize_matplotlib
import numpy as np

(x_train, y_train), (x_test, y_test) = keras.datasets.cifar10.load_data()
x_train = x_train.astype('float32') / 255.0
x_test  = x_test.astype('float32')  / 255.0
y_test_flat = y_test.flatten()

def build_model():
    model = keras.Sequential([
        keras.layers.Input(shape=(32, 32, 3)),
        keras.layers.Conv2D(64, (3, 3), activation='relu', padding='same'),
        keras.layers.MaxPooling2D((2, 2)),
        keras.layers.Conv2D(128, (3, 3), activation='relu', padding='same'),
        keras.layers.MaxPooling2D((2, 2)),
        keras.layers.GlobalAveragePooling2D(),
        keras.layers.Dense(128, activation='relu'),
        keras.layers.Dropout(0.2),
        keras.layers.Dense(10, activation='softmax'),
    ])
    return model

def compile_and_fit(use_ema):
    keras.backend.clear_session()
    model = build_model()
    optimizer = keras.optimizers.Adam(
        learning_rate=0.001,
        use_ema=use_ema,
        ema_momentum=0.99 if use_ema else None,
    )
    model.compile(optimizer=optimizer,
                  loss='sparse_categorical_crossentropy',
                  metrics=['accuracy'])
    # model.fit()終了時にEMA重みへ自動的に切り替わる(finalize_variable_values()の手動呼び出しは不要)
    history = model.fit(x_train, y_train, epochs=30, batch_size=64,
                        validation_split=0.2, verbose=1)
    return model, history
実行結果をクリックして内容を開く
Downloading data from https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz
170498071/170498071 ━━━━━━━━━━━━━━━━━━━━ 1591s 9us/step

TTA予測関数

def predict_with_tta(model, x):
    # 元画像と水平反転画像、2種類の予測確率を平均する
    probs_original = model.predict(x, verbose=0)
    x_flipped = x[:, :, ::-1, :]
    probs_flipped = model.predict(x_flipped, verbose=0)
    return (probs_original + probs_flipped) / 2.0

def evaluate_accuracy(probs, y_true):
    preds = np.argmax(probs, axis=1)
    return float(np.mean(preds == y_true))
実行結果をクリックして内容を開く
出力なし

4パターンの学習・評価実行

ema_configs = [(False, 'no_ema'), (True, 'with_ema')]
models, histories, base_probs = {}, {}, {}

for use_ema, key in ema_configs:
    print(f"\n=== 学習: use_ema={use_ema} ===")
    model, h = compile_and_fit(use_ema)
    models[key] = model
    histories[key] = h
    # TTAなしの予測確率(通常のevaluateと同じ結果になるはず)
    base_probs[key] = model.predict(x_test, verbose=0)

results = {}
for use_ema, key in ema_configs:
    model = models[key]

    # TTAなし
    acc_no_tta = evaluate_accuracy(base_probs[key], y_test_flat)
    # TTAあり
    tta_probs = predict_with_tta(model, x_test)
    acc_with_tta = evaluate_accuracy(tta_probs, y_test_flat)

    results[key] = {'no_tta': acc_no_tta, 'with_tta': acc_with_tta}
    print(f"{key}: TTAなし={acc_no_tta:.4f}  TTAあり={acc_with_tta:.4f}")
実行結果をクリックして内容を開く
=== 学習: use_ema=False ===
Epoch 1/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 17s 12ms/step - accuracy: 0.2646 - loss: 1.9331 - val_accuracy: 0.3621 - val_loss: 1.7178
Epoch 2/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 11s 7ms/step - accuracy: 0.3718 - loss: 1.6744 - val_accuracy: 0.4275 - val_loss: 1.5595
Epoch 3/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4246 - loss: 1.5574 - val_accuracy: 0.4532 - val_loss: 1.4862
Epoch 4/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4659 - loss: 1.4548 - val_accuracy: 0.4647 - val_loss: 1.4611
Epoch 5/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.4924 - loss: 1.3855 - val_accuracy: 0.5186 - val_loss: 1.3134
Epoch 6/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5117 - loss: 1.3379 - val_accuracy: 0.5371 - val_loss: 1.2826
Epoch 7/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5231 - loss: 1.3004 - val_accuracy: 0.5464 - val_loss: 1.2451
Epoch 8/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5421 - loss: 1.2613 - val_accuracy: 0.5539 - val_loss: 1.2118
Epoch 9/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.5548 - loss: 1.2318 - val_accuracy: 0.5610 - val_loss: 1.2252
Epoch 10/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5629 - loss: 1.2028 - val_accuracy: 0.5670 - val_loss: 1.1888
Epoch 11/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5741 - loss: 1.1768 - val_accuracy: 0.5751 - val_loss: 1.1545
Epoch 12/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5821 - loss: 1.1560 - val_accuracy: 0.5858 - val_loss: 1.1417
Epoch 13/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5918 - loss: 1.1337 - val_accuracy: 0.5862 - val_loss: 1.1439
Epoch 14/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 8ms/step - accuracy: 0.5961 - loss: 1.1164 - val_accuracy: 0.6002 - val_loss: 1.1030
Epoch 15/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6077 - loss: 1.0895 - val_accuracy: 0.5938 - val_loss: 1.1044
Epoch 16/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6177 - loss: 1.0696 - val_accuracy: 0.6158 - val_loss: 1.0605
Epoch 17/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6215 - loss: 1.0606 - val_accuracy: 0.6104 - val_loss: 1.0648
Epoch 18/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6288 - loss: 1.0348 - val_accuracy: 0.6110 - val_loss: 1.0598
Epoch 19/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6309 - loss: 1.0268 - val_accuracy: 0.6178 - val_loss: 1.0675
Epoch 20/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6395 - loss: 1.0066 - val_accuracy: 0.6299 - val_loss: 1.0214
Epoch 21/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 7ms/step - accuracy: 0.6451 - loss: 0.9937 - val_accuracy: 0.6472 - val_loss: 0.9855
Epoch 22/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6500 - loss: 0.9754 - val_accuracy: 0.6525 - val_loss: 0.9733
Epoch 23/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6546 - loss: 0.9677 - val_accuracy: 0.6545 - val_loss: 0.9600
Epoch 24/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6595 - loss: 0.9500 - val_accuracy: 0.6625 - val_loss: 0.9435
Epoch 25/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6669 - loss: 0.9395 - val_accuracy: 0.6563 - val_loss: 0.9528
Epoch 26/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6678 - loss: 0.9290 - val_accuracy: 0.6749 - val_loss: 0.9165
Epoch 27/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6722 - loss: 0.9207 - val_accuracy: 0.6648 - val_loss: 0.9340
Epoch 28/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6775 - loss: 0.9021 - val_accuracy: 0.6701 - val_loss: 0.9193
Epoch 29/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6794 - loss: 0.8938 - val_accuracy: 0.6623 - val_loss: 0.9365
Epoch 30/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6837 - loss: 0.8878 - val_accuracy: 0.6542 - val_loss: 0.9551

=== 学習: use_ema=True ===
Epoch 1/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 9s 9ms/step - accuracy: 0.2555 - loss: 1.9410 - val_accuracy: 0.3380 - val_loss: 1.7297
Epoch 2/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.3638 - loss: 1.6881 - val_accuracy: 0.3943 - val_loss: 1.6180
Epoch 3/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4279 - loss: 1.5557 - val_accuracy: 0.4499 - val_loss: 1.5142
Epoch 4/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 6s 7ms/step - accuracy: 0.4680 - loss: 1.4628 - val_accuracy: 0.4859 - val_loss: 1.3994
Epoch 5/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4855 - loss: 1.4031 - val_accuracy: 0.5003 - val_loss: 1.3581
Epoch 6/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5013 - loss: 1.3570 - val_accuracy: 0.5205 - val_loss: 1.3192
Epoch 7/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5172 - loss: 1.3224 - val_accuracy: 0.5169 - val_loss: 1.3260
Epoch 8/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5302 - loss: 1.2887 - val_accuracy: 0.5301 - val_loss: 1.2725
Epoch 9/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5439 - loss: 1.2552 - val_accuracy: 0.5562 - val_loss: 1.2069
Epoch 10/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5540 - loss: 1.2276 - val_accuracy: 0.5654 - val_loss: 1.1891
Epoch 11/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5627 - loss: 1.2034 - val_accuracy: 0.5794 - val_loss: 1.1621
Epoch 12/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5688 - loss: 1.1847 - val_accuracy: 0.5806 - val_loss: 1.1530
Epoch 13/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 6s 7ms/step - accuracy: 0.5793 - loss: 1.1578 - val_accuracy: 0.5858 - val_loss: 1.1301
Epoch 14/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5872 - loss: 1.1396 - val_accuracy: 0.5971 - val_loss: 1.1102
Epoch 15/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5933 - loss: 1.1238 - val_accuracy: 0.5998 - val_loss: 1.1104
Epoch 16/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5983 - loss: 1.1076 - val_accuracy: 0.6090 - val_loss: 1.0765
Epoch 17/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6084 - loss: 1.0823 - val_accuracy: 0.6106 - val_loss: 1.0635
Epoch 18/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6153 - loss: 1.0683 - val_accuracy: 0.6122 - val_loss: 1.0575
Epoch 19/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6206 - loss: 1.0533 - val_accuracy: 0.6167 - val_loss: 1.0602
Epoch 20/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.6275 - loss: 1.0346 - val_accuracy: 0.6360 - val_loss: 1.0188
Epoch 21/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6323 - loss: 1.0211 - val_accuracy: 0.6275 - val_loss: 1.0217
Epoch 22/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6389 - loss: 1.0037 - val_accuracy: 0.6462 - val_loss: 0.9832
Epoch 23/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6479 - loss: 0.9925 - val_accuracy: 0.6291 - val_loss: 1.0333
Epoch 24/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.6488 - loss: 0.9786 - val_accuracy: 0.6473 - val_loss: 0.9802
Epoch 25/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6531 - loss: 0.9678 - val_accuracy: 0.6511 - val_loss: 0.9617
Epoch 26/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6554 - loss: 0.9551 - val_accuracy: 0.6633 - val_loss: 0.9500
Epoch 27/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 7ms/step - accuracy: 0.6614 - loss: 0.9408 - val_accuracy: 0.6563 - val_loss: 0.9497
Epoch 28/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 7ms/step - accuracy: 0.6645 - loss: 0.9364 - val_accuracy: 0.6667 - val_loss: 0.9372
Epoch 29/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6699 - loss: 0.9200 - val_accuracy: 0.6627 - val_loss: 0.9332
Epoch 30/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6726 - loss: 0.9093 - val_accuracy: 0.6559 - val_loss: 0.9438
no_ema: TTAなし=0.6550  TTAあり=0.6599
with_ema: TTAなし=0.6838  TTAあり=0.6869

グラフ+サマリー

# ── val_accuracy 比較グラフ ───────────────
fig, ax = plt.subplots(figsize=(8, 5))
for key, h in histories.items():
    ax.plot(h.history['val_accuracy'], label=key)
ax.set_title('val_accuracy の比較(全30エポック)')
ax.set_xlabel('Epoch'); ax.legend(); ax.grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig('tta_ema_comparison.png', dpi=150)
plt.show()

print("\n===== 最終結果サマリー =====")
print(f"{'Pattern':>12} | {'EMA':>5} | {'TTA':>5} | {'Test Acc':>9}")
print("-" * 42)
print(f"{'A_baseline':>12} | {'なし':>5} | {'なし':>5} | {results['no_ema']['no_tta']:>9.4f}")
print(f"{'B_ema_only':>12} | {'あり':>5} | {'なし':>5} | {results['with_ema']['no_tta']:>9.4f}")
print(f"{'C_tta_only':>12} | {'なし':>5} | {'あり':>5} | {results['no_ema']['with_tta']:>9.4f}")
print(f"{'D_ema_tta':>12} | {'あり':>5} | {'あり':>5} | {results['with_ema']['with_tta']:>9.4f}")
print("-" * 42)
実行結果をクリックして内容を開く
===== 最終結果サマリー =====
     Pattern |   EMA |   TTA |  Test Acc
------------------------------------------
  A_baseline |    なし |    なし |    0.6550
  B_ema_only |    あり |    なし |    0.6838
  C_tta_only |    なし |    あり |    0.6599
   D_ema_tta |    あり |    あり |    0.6869
------------------------------------------

実験結果

精度グラフ

精度グラフ

結果サマリー

パターンEMATTAtest_accuracybaseline比
A:baselineなしなし65.50%±0
B:EMAのみありなし68.38%+2.88pt
C:TTAのみなしあり65.99%+0.49pt
D:EMA+TTAありあり68.69%+3.19pt

単体の改善幅を単純合計すると+3.37pt(68.87%)になるところ、実測のDは+3.19pt(68.69%)とほぼ近い水準でした。わずかに合計を下回るものの、大きな重複や頭打ちは見られませんでした。

ハマりポイント

  • EMAはoptimizer作成時にuse_ema=Trueを指定するだけでよく、model.fit()終了時に自動的にEMA重みへ切り替わる。過去のKerasの使い方に慣れているとoptimizer.finalize_variable_values()を手動で呼びたくなるが、fit()を最後まで実行していれば不要(二重に呼んでもエラーにはならないが冗長)。
  • 水平反転TTAはx[:, :, ::-1, :]のようにNumPyのスライスで実装できるが、対象軸を間違えやすい。CIFAR-10の画像は(N, H, W, C)形式なので、反転すべきは幅方向(3番目の軸、インデックス2)であり、高さ方向(インデックス1)を反転すると天地が逆になった不自然な画像になってしまう。
  • TTAの平均化はsoftmax後の確率空間で行う。softmax前のロジットのまま平均してからargmaxを取ると、確率空間での平均と数学的に一致しない(softmaxは非線形なため)。model.predict()はデフォルトでsoftmax適用後の確率を返すため、今回のコードはそのまま平均すればよい。
  • D(EMA+TTA)の効果を正しく評価するには、Bで学習したEMA適用済みモデルに対してTTAを適用する必要がある。EMAなしで学習したモデルにTTAだけ追加しても、それはCと同じ実験になってしまう。

考察

① EMA単体・TTA単体はそれぞれどの程度効果があったか

EMA単体は+2.88pt(65.50%→68.38%)、TTA単体は+0.49pt(65.50%→65.99%)と、EMAの効果はTTAの約6倍という大差がつきました。これは「作用する場所の違い」で説明できます。EMAは学習トラジェクトリ全体(重みの推移)に対して働くため、学習終盤のミニバッチごとのブレを面で均す、いわば大域的な補正です。一方、今回のTTAは水平反転による2-view平均に留まっており、決定境界付近の際どいサンプルの予測だけを揺らす局所的な補正にとどまります。TTAのview数を増やす(クロップやズームを追加する)ことで、この差がどこまで縮むかは別途検証の余地があります。

② 組み合わせ(D)は単体の改善幅の合計に近いか、それとも頭打ちか

単体の改善幅を単純合計すると+3.37pt(68.87%)ですが、実測のDは+3.19pt(68.69%)で、0.18pt分だけ合計に届いていません。ここで「ほぼ合計通りだから加算的」と結論づけるより、TTAの限界効果(marginal effect)で見る方が実態を捉えられます。

  • TTAをbaselineに追加した場合の効果:+0.49pt(A→C)
  • TTAをEMA適用済みモデルに追加した場合の効果:+0.31pt(B→D)

同じ「水平反転TTAを追加する」という操作でも、EMAが既に効いているモデルに対しては効果が約37%目減りしています。これは、EMAが学習終盤のノイズをすでにある程度均してしまっているため、TTAが平均化できる「残された分散」がEMAなしの場合より少なくなっている、と解釈するのが自然です。つまり完全に独立(加算的)でも完全に重複(頭打ち)でもなく、「EMAがTTAの効きしろの一部を先取りしている」という部分的重複が今回の結果です。

③ 重み空間と予測空間、どちらの分散低減がより効果的だったか

今回の設定では重み空間(EMA)の分散低減の方が予測空間(TTA)より圧倒的に効果的でした(+2.88pt vs +0.49pt)。ただしこれは「EMAが常に優れている」という一般論ではなく、TTAが2-viewという最小構成だったことの影響も大きいと考えられます。学習コストをほぼ増やさずに効くEMAに対し、TTAは推論コストを増やしてなお単体では控えめな効果に留まりました。この非対称性は、限られたリソースでどちらを優先すべきかを考える上で重要な結果です。

④ 想定との差分と確認ポイント

今回の結果は「大きく重複する」でも「完全に独立」でもなく、EMAがTTAの効きしろの一部(約37%)を先取りする、部分的な重複という、事前に用意していた2択(加算的/頭打ち)のどちらとも言い切れない中間的な結果でした。もし手元の実験でDがBを下回る(TTAを足すとむしろ悪化する)ような、より明確な頭打ちが見られた場合は、以下を確認してください。

  • ema_momentumの値が適切か(今回は0.99。小さすぎるとEMAの平滑化効果がほぼ効かない)
  • TTAの反転処理が正しい軸に対して行われているか(画像を目視で確認する)
  • D(EMA+TTA)が、Bで学習したEMA適用済みモデルに対して正しくTTAを適用できているか(BとDで同じモデルを使っているか確認する)
ケース実務での推奨
学習コストを増やさず精度を上げたいEMAが最有力(+2.88pt)。use_ema=Trueを指定するだけで学習コストはほぼ増えないため、まず試す価値が高い。
推論コストを増やさず精度を上げたいEMAを推奨。TTAと違い推論時のコスト増加がない。
推論時間に余裕があり、さらに精度を積みたいEMA適用済みモデルにTTAを追加すると+0.31pt程度の上乗せが見込める(単体効果+0.49ptよりは目減りする点に留意)。
両方使う場合のコスト感EMAは学習コストの増加はほぼゼロ、TTAは推論時間がview数倍になる点に注意。今回は2-view(原画像+反転)なので推論時間は約2倍になる。

まとめ

EMA単体(+2.88pt)はTTA単体(+0.49pt)の約6倍の効果があり、組み合わせ(+3.19pt)は単純合計(+3.37pt)にほぼ近い水準でした。ただし限界効果で見ると、TTAの効きしろはEMAによって約37%先取りされており(単体+0.49pt→EMA併用時+0.31pt)、完全な独立でも頭打ちでもない「部分的な重複」が今回の結論です。学習コストがほぼゼロで効果の大きいEMAをまず適用し、推論コストに余裕があればTTAを追加する、という優先順位が今回の結果からは妥当と言えます。


関連記事もあわせてどうぞ(URLは公開時に要確認):