Test Time Augmentation(TTA)で精度は上がる?あり/なし比較【Keras×CIFAR-10実験】

投稿日:2026年7月24日金曜日 最終更新日:

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

X f B! L
Test Time Augmentation(TTA)で精度は上がる?あり/なし比較【Keras×CIFAR-10実験】 アイキャッチ画像

これまでSWA・EMAと、学習後の重みを平均するテクニックを2本続けて試してきました。今回はアプローチを変えて、推論時の予測を平均するTest Time Augmentation(TTA)を試します。

TTAは、1枚のテスト画像をそのまま推論するのではなく、水平反転やシフトなど複数のバリエーション(View)を作って個別に推論し、その予測確率を平均して最終的な予測とする手法です。SWA・EMAとの決定的な違いは再学習が一切不要という点です。すでに学習済みのモデルに対して、推論時の処理を変えるだけで試せます。ただし過去の記事で「CIFAR-10ではflip以外の複雑なAugmentationは効きにくい」という結果が出ているため、TTAのViewを増やせば増やすほど良いとは限りません。今回はView数を変えた2パターンで検証します。

📘 この記事でわかること
  • Test Time Augmentation(TTA)の仕組みと、SWA・EMAとの違い
  • Kerasでflip・shiftを使ったTTAを実装する方法
  • Viewの数(2 view vs 6 view)で精度・推論時間がどう変わるか
  • 学習時に使っていないAugmentationをTTAに使うとどうなるか

TTAとは何をしているのか

通常の推論では、1枚のテスト画像 $x$ をそのままモデルに入力し、予測確率 $p = f(x)$ を得ます。TTAでは、$x$ に対して $n$ 種類のView変換 $t_1, \dots, t_n$ を適用し、それぞれの予測確率を平均します。

$$ p_{TTA} = \frac{1}{n} \sum_{i=1}^{n} f(t_i(x)) $$

学習時のData Augmentationが「学習データのバリエーションを増やしてモデルを頑健にする」ものだとすると、TTAは「推論データのバリエーションを増やして予測を安定させる」ものといえます。複数の見え方に対する予測を平均することで、1回の推論だけではブレてしまう境界線上のサンプルの判定が安定する、という考え方です。

項目SWA/EMATTA
平均する対象モデルの重み推論時の予測確率
再学習の要否学習時に組み込む必要がある不要(学習済みモデルにそのまま適用可)
推論コスト通常と同じ(1回のforward)View数倍のforwardが必要
Kerasでの対応EMAはOptimizerに公式実装あり公式実装なし。自前で推論ループを書く必要がある

実験コード

使用環境はGoogle Colab(GPU:T4)、データセットはCIFAR-10です。ベースラインモデル(Conv2D×2層、GAP、Dropout=0.2、Adam lr=0.001、batch_size=64、30エポック、Data Augmentationなし)で1回だけ学習し、推論時のTTAパターンだけを変えて比較します。SWA・EMAの記事とは異なり、学習自体は1回で済むのが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,721 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 61.6 MB/s eta 0:00:00
  Preparing metadata (setup.py) ... done
  Building wheel for japanize_matplotlib (setup.py) ... done
環境準備完了

import・データ準備・モデル構築・学習(1回のみ)

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

(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 set_seed(seed=42):
    random.seed(seed)
    np.random.seed(seed)
    tf.random.set_seed(seed)

set_seed(42)

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'),
])
model.compile(optimizer='adam',
              loss='sparse_categorical_crossentropy',
              metrics=['accuracy'])

history = model.fit(x_train, y_train, epochs=30, batch_size=64,
                    validation_split=0.2, verbose=1)
実行結果をクリックして内容を開く
Downloading data from https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz
170498071/170498071 ━━━━━━━━━━━━━━━━━━━━ 3158s 19us/step
Epoch 1/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 10s 8ms/step - accuracy: 0.2577 - loss: 1.9335 - val_accuracy: 0.3526 - val_loss: 1.7267
Epoch 2/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.3539 - loss: 1.7018 - val_accuracy: 0.4238 - val_loss: 1.5883
Epoch 3/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.4203 - loss: 1.5785 - val_accuracy: 0.4707 - val_loss: 1.4753
Epoch 4/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4593 - loss: 1.4802 - val_accuracy: 0.4934 - val_loss: 1.3977
Epoch 5/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4884 - loss: 1.4021 - val_accuracy: 0.5083 - val_loss: 1.3466
Epoch 6/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5047 - loss: 1.3531 - val_accuracy: 0.5209 - val_loss: 1.3028
Epoch 7/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5231 - loss: 1.3120 - val_accuracy: 0.5387 - val_loss: 1.2651
Epoch 8/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5355 - loss: 1.2763 - val_accuracy: 0.5424 - val_loss: 1.2464
Epoch 9/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5506 - loss: 1.2421 - val_accuracy: 0.5592 - val_loss: 1.2078
Epoch 10/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5586 - loss: 1.2117 - val_accuracy: 0.5625 - val_loss: 1.1937
Epoch 11/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5702 - loss: 1.1823 - val_accuracy: 0.5712 - val_loss: 1.1631
Epoch 12/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5818 - loss: 1.1585 - val_accuracy: 0.5765 - val_loss: 1.1638
Epoch 13/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5888 - loss: 1.1371 - val_accuracy: 0.5888 - val_loss: 1.1293
Epoch 14/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6021 - loss: 1.1110 - val_accuracy: 0.6037 - val_loss: 1.0901
Epoch 15/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6036 - loss: 1.0942 - val_accuracy: 0.6031 - val_loss: 1.0805
Epoch 16/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6114 - loss: 1.0726 - val_accuracy: 0.6133 - val_loss: 1.0624
Epoch 17/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6209 - loss: 1.0564 - val_accuracy: 0.6223 - val_loss: 1.0379
Epoch 18/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6280 - loss: 1.0401 - val_accuracy: 0.6290 - val_loss: 1.0180
Epoch 19/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.6329 - loss: 1.0234 - val_accuracy: 0.6337 - val_loss: 1.0029
Epoch 20/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6394 - loss: 1.0074 - val_accuracy: 0.6338 - val_loss: 1.0110
Epoch 21/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6460 - loss: 0.9909 - val_accuracy: 0.6422 - val_loss: 0.9885
Epoch 22/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6505 - loss: 0.9796 - val_accuracy: 0.6498 - val_loss: 0.9740
Epoch 23/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6527 - loss: 0.9666 - val_accuracy: 0.6554 - val_loss: 0.9612
Epoch 24/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6611 - loss: 0.9515 - val_accuracy: 0.6559 - val_loss: 0.9555
Epoch 25/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6648 - loss: 0.9392 - val_accuracy: 0.6579 - val_loss: 0.9518
Epoch 26/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6687 - loss: 0.9278 - val_accuracy: 0.6591 - val_loss: 0.9530
Epoch 27/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6726 - loss: 0.9141 - val_accuracy: 0.6645 - val_loss: 0.9373
Epoch 28/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6750 - loss: 0.9043 - val_accuracy: 0.6705 - val_loss: 0.9241
Epoch 29/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6795 - loss: 0.8921 - val_accuracy: 0.6708 - val_loss: 0.9190
Epoch 30/30
625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6883 - loss: 0.8767 - val_accuracy: 0.6695 - val_loss: 0.9276

TTAの実装と3パターンの評価

# ── TTA用のView定義 ───────────────────────────────
# Kerasに公式のTTA実装はないため、Viewごとに推論し確率を平均する処理を自作する。
def view_identity(x):
    return x

def view_hflip(x):
    return x[:, :, ::-1, :]  # 幅方向(左右)を反転

def make_shift_view(dy, dx, pad=4):
    # reflectパディングしてから (dy, dx) だけずらして元のサイズに切り出す
    def _shift(x):
        padded = np.pad(x, ((0, 0), (pad, pad), (pad, pad), (0, 0)), mode='reflect')
        return padded[:, pad + dy: pad + dy + 32, pad + dx: pad + dx + 32, :]
    return _shift


def predict_with_tta(model, x, views, batch_size=256):
    probs_sum = None
    for view_fn in views:
        x_view = view_fn(x)
        probs = model.predict(x_view, batch_size=batch_size, verbose=0)
        probs_sum = probs if probs_sum is None else probs_sum + probs
    return probs_sum / len(views)


views_config = [
    ("A_no_tta", [view_identity]),
    ("B_tta_flip2", [view_identity, view_hflip]),
    ("C_tta_flip_shift6", [
        view_identity, view_hflip,
        make_shift_view(-4, -4), make_shift_view(-4, 4),
        make_shift_view(4, -4),  make_shift_view(4, 4),
    ]),
]

results = {}
for name, views in views_config:
    label = name.split('_', 1)[1]
    start = time.time()
    avg_probs = predict_with_tta(model, x_test, views)
    elapsed = time.time() - start
    preds = np.argmax(avg_probs, axis=1)
    acc = np.mean(preds == y_test_flat)
    results[label] = {"accuracy": acc, "time": elapsed, "n_views": len(views)}
    print(f"[{label}] View数:{len(views)} test_accuracy:{acc:.4f} 推論時間:{elapsed:.2f}秒")
実行結果をクリックして内容を開く
[no_tta] View数:1 test_accuracy:0.6701 推論時間:1.79秒
[tta_flip2] View数:2 test_accuracy:0.6727 推論時間:1.16秒
[tta_flip_shift6] View数:6 test_accuracy:0.6604 推論時間:3.82秒

グラフ+サマリー

# ── 学習曲線+TTAパターン別の精度・推論時間を可視化 ─────
fig, axes = plt.subplots(1, 3, figsize=(18, 5))

axes[0].plot(history.history['val_accuracy'], label='val_accuracy')
axes[0].plot(history.history['accuracy'], label='accuracy')
axes[0].set_title('学習曲線(モデルは1つのみ)')
axes[0].set_xlabel('Epoch'); axes[0].legend(); axes[0].grid(True, alpha=0.3)

labels = list(results.keys())
accs = [results[l]["accuracy"] for l in labels]
times_ = [results[l]["time"] for l in labels]

axes[1].bar(labels, accs, color=['gray', 'steelblue', 'darkorange'])
axes[1].set_title('TTAパターン別 test_accuracy')
axes[1].set_ylim(min(accs) - 0.02, max(accs) + 0.02)
axes[1].tick_params(axis='x', rotation=20)

axes[2].bar(labels, times_, color=['gray', 'steelblue', 'darkorange'])
axes[2].set_title('TTAパターン別 推論時間(秒)')
axes[2].tick_params(axis='x', rotation=20)

plt.tight_layout()
plt.savefig('tta_comparison.png', dpi=150)
plt.show()

print("\n===== 最終結果サマリー =====")
print(f"{'Pattern':>18} | {'View数':>6} | {'test_accuracy':>13} | {'推論時間(s)':>10}")
print("-" * 58)
for label in ['no_tta', 'tta_flip2', 'tta_flip_shift6']:
    r = results[label]
    print(f"{label:>18} | {r['n_views']:>6} | {r['accuracy']:>13.4f} | {r['time']:>10.2f}")
print("-" * 58)

最終結果サマリー

===== 最終結果サマリー =====
           Pattern |  View数 | test_accuracy |    推論時間(s)
----------------------------------------------------------
            no_tta |      1 |        0.6701 |       1.79
         tta_flip2 |      2 |        0.6727 |       1.16
   tta_flip_shift6 |      6 |        0.6604 |       3.82
----------------------------------------------------------

実験結果

学習曲線・パターン別比較グラフ

学習曲線

学習曲線

TTAパターン別 test_accuracy

TTAパターン別 test_accuracy

TTAパターン別 推論時間(秒)

TTAパターン別 推論時間(秒)
パターンView数test_accuracy差分(vs TTAなし)推論時間(実測)
A:TTAなし167.01%1.79秒
B:TTAあり(flip、2 view)267.27%(+0.26pt)+0.26pt1.16秒
C:TTAあり(flip+shift、6 view)666.04%(−0.97pt)−0.97pt3.82秒

Aの推論時間(1.79秒)は、View数に対して線形に増える関係(B・Cから算出:1 viewあたり約0.665秒、1 viewなら理論上約0.5秒)から見て不自然に長い値です。これはTTAの仕組みとは無関係な測定上の癖である可能性が高く、次のハマりポイントで検証します。

⚠ ハマりポイント
  • Viewの変換は学習時のAugmentationと揃える必要はないが、揃っていない場合は効果が読みにくい:今回のベースラインはData Augmentationなしで学習しているため、モデルはshiftされた画像を一度も見ていない。TTAでshift Viewを追加した場合、学習分布から外れた入力に対する予測を平均に混ぜることになり、精度に悪影響を与える可能性がある。
  • flip Viewはx[:, :, ::-1, :]でNumPy配列のまま反転できる:わざわざtf.image.flip_left_rightを使わなくても、軸方向のスライス反転で十分。バッチ全体に対して一括で処理できるため高速。
  • Viewを増やすほど推論時間はほぼ線形に増える:6 viewなら1 viewの場合の約6倍のforward計算が必要になる。学習コストはゼロでも、推論コスト(特にリアルタイム性が求められる用途)には直結する点に注意。
  • ループの最初のmodel.predict()呼び出しは、GPU・cuDNNのカーネル初期化などのウォームアップ分だけ余計に時間がかかる:今回の実測でも、View数1(TTAなし)が最も少ないView数にもかかわらず、View数2(flip)より長い時間がかかった。B・Cの実測値から算出した線形関係(1 viewあたり約0.665秒)で推定すると、View数1は本来約0.5秒のはずが、実測は1.79秒だった。差分の約1.3秒は、ループの1回目に発生したウォームアップコストと考えられる。推論時間を正確に比較したい場合は、計測開始前に一度ダミーのpredict()を実行してウォームアップを済ませておくべきだった。

考察

① TTA(flipのみ、2 view)は精度向上に寄与したか

test_accuracyはTTAなしの67.01%からflip TTA(2 view)で67.27%へと、+0.26ptのわずかな改善が見られました。再学習なしで、推論をView数分繰り返すだけでこの程度の改善が得られるのは、コストパフォーマンスとしては悪くありません。ただし改善幅自体は小さく、ノイズの範囲に近いことは踏まえておく必要があります。

② Viewを増やす(flip+shift、6 view)と結果はどう変わったか

flip+shiftの6 viewでは66.04%となり、TTAなし(67.01%)よりも−0.97pt、flipのみの2 view(67.27%)よりも−1.23pt低下しました。これは過去の記事「Data Augmentationを重ねすぎると精度が下がる?」で確認した、flip以外の複雑なAugmentationはCIFAR-10では効きにくいという傾向と一致する結果です。

今回のベースラインモデルはshift(平行移動)を含むAugmentationなしで学習しているため、モデルはshiftされた画像分布を見たことがありません。TTAでshift Viewを追加すると、モデルが不慣れな入力に対する(信頼度の低い)予測が平均に混ざり込み、むしろ予測全体の質を下げてしまったと考えられます。TTAのViewは、学習時に慣れ親しんだ変換の種類を選ぶことが重要という結論になります。

③ 推論時間はView数に比例して増えたか——ただし測定には注意が必要

B(2 view:1.16秒)とC(6 view:3.82秒)の間ではおおむね線形の関係(View数4個の増加に対し+2.66秒、1 viewあたり約0.665秒)が確認できました。一方でA(1 view:1.79秒)はこの関係から外れており、線形関係から予測される値(約0.5秒)より大幅に長い結果でした。これはハマりポイントで触れた通り、ループの最初の推論呼び出しに伴うGPUのウォームアップコストが混入したためと考えられます。View数と推論時間の関係自体はほぼ線形とみてよいですが、絶対値を比較する際はウォームアップの影響に注意が必要です。

実務での推奨

状況推奨
学習済みモデルの精度をもう少し上げたいが、再学習はしたくないTTAは有力な選択肢。特にflipのような軽量なViewから試す
Viewの種類を選ぶ際学習時に使ったAugmentationと同じ種類のViewを選ぶと、モデルが見慣れた入力分布の範囲内で予測できる
推論速度・リアルタイム性が重要な用途View数を増やすほど推論時間が線形に増えるため、精度向上とのトレードオフを見て判断する
バッチ推論・オフライン評価推論時間の制約が緩いため、View数を増やしたTTAを積極的に試す価値がある
✅ まとめ
  • TTAは推論時に複数のViewの予測確率を平均する手法で、SWA・EMAと異なり再学習が不要
  • Kerasに公式実装はなく、Viewごとにmodel.predict()を呼んで確率を平均する処理を自前で書く必要がある
  • 今回の実験では、flipのみ(2 view)は+0.26ptの改善flip+shift(6 view)は−0.97ptの悪化となった。学習時に使っていない変換(shift)をTTAに混ぜると逆効果になりやすい
  • Viewの数を増やすほど推論時間はほぼ線形に増える。ただし計測時はGPUのウォームアップの影響に注意(1回目の推論呼び出しだけ余計に時間がかかることがある)

関連記事もあわせてどうぞ: