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ではoptimizerにuse_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 入力)が異なるため、理屈の上では独立に効果を持ちやすいはずですが、両方とも最終的には「モデルの予測のブレを減らす」という同じ目的地に向かっているため、重複が生じる可能性もあります。
| パターン | EMA | TTA | 平均化の対象 |
|---|---|---|---|
| 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
------------------------------------------
実験結果
精度グラフ
結果サマリー
| パターン | EMA | TTA | test_accuracy | baseline比 |
|---|---|---|---|---|
| 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は公開時に要確認):


0 件のコメント:
コメントを投稿