前回の記事では、学習後半の重みを単純平均するSWA(Stochastic Weight Averaging)を試したところ、モデルがまだ収束していない状態で平均を取ってしまい、むしろ精度が下がるという結果になりました。
今回試すEMA(Exponential Moving Average:指数移動平均)は、SWAと同じ「重みを平均する」系統のテクニックですが、仕組みが異なります。直近の重みほど強く反映し、古い重みの影響は指数的に減衰していきます。SWAが苦手だった「まだ収束していない学習」でも、EMAなら通用するのでしょうか。KerasのOptimizerにはuse_emaという引数で公式にEMAのサポートがあるため、今回はそれを使って検証します。
- EMA(指数移動平均)の仕組みと、SWAとの違い
- KerasのOptimizerに組み込まれた
use_ema/ema_momentumの使い方 ema_momentum(減衰率)の値によって結果がどう変わるか- EMAを使う上でハマりやすいポイント(
finalize_variable_values()の必要性など)
EMAとは何をしているのか、SWAとの違い
EMAは、学習の各ステップ(バッチ)ごとに、それまでのEMA重みと現在の重みを一定の比率で混ぜ合わせていきます。
$$ \theta_{EMA}^{(t)} = \text{decay} \times \theta_{EMA}^{(t-1)} + (1 - \text{decay}) \times \theta^{(t)} $$
$\text{decay}$(Kerasではema_momentum)が1に近いほど、古い重みの影響が長く残ります。SWAが「指定した区間の重みを均等に平均する」のに対し、EMAは「区間を明示的に指定せず、直近ほど重みを強くしながら学習全体を通して平均し続ける」という違いがあります。
| 項目 | SWA | EMA |
|---|---|---|
| 重みの付け方 | 指定区間内で均等 | 指数的に減衰(直近ほど重視) |
| 適用範囲 | 通常は学習後半のみ(開始epochを指定) | 学習全体を通して更新可能 |
| Kerasでの対応 | 公式実装なし。自作コールバックが必要 | Optimizerのuse_ema=Trueで公式サポート |
| 平均後の重みの反映 | set_weights()で即座に反映 | fit()使用時は最終epoch後に自動反映(finalize_variable_values()は本来手動制御用) |
もう一点、SWAの記事では触れなかった重要な注意点があります。EMAは学習の最初のステップから適用されるため、$\text{decay}$の値によっては「まだ何も学習していない初期のランダムな重み」の影響がいつまでも残ってしまうことがあります。この点は後の考察で実際の数値とともに検証します。
実験コード
使用環境はGoogle Colab(GPU:T4)、データセットはCIFAR-10です。ベースラインモデル(Conv2D×2層、GAP、Dropout=0.2、Adam lr=0.001、batch_size=64、30エポック)はSWAの記事と同一で、EMAの有無とema_momentumの値だけを変えて比較します。
環境準備(最初に一度だけ実行)
# ── 環境準備(最初に一度だけ実行)──────────────────────
!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,231 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 56.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 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
def set_seed(seed=42):
random.seed(seed)
np.random.seed(seed)
tf.random.set_seed(seed)
def build_model(name):
return 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'),
], name=name)
実行結果をクリックして内容を開く
Downloading data from https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz 170498071/170498071 ━━━━━━━━━━━━━━━━━━━━ 2584s 15us/step
3パターンの学習実行(Optimizer組み込みのEMAを使用)
# ── Optimizer組み込みのEMAを使う ────────────────────
# KerasのOptimizerはuse_ema=Trueとema_momentumで
# 指数移動平均を公式にサポートしている(自作コールバック不要)。
def compile_and_fit(model, ema_momentum=None):
use_ema = ema_momentum is not None
optimizer = keras.optimizers.Adam(
learning_rate=0.001,
use_ema=use_ema,
ema_momentum=ema_momentum if use_ema else 0.99, # use_ema=Falseの場合は無視される
)
model.compile(optimizer=optimizer,
loss='sparse_categorical_crossentropy',
metrics=['accuracy'])
start = time.time()
history = model.fit(x_train, y_train, epochs=30, batch_size=64,
validation_split=0.2, verbose=1)
elapsed = time.time() - start
return history, elapsed, optimizer, use_ema
configs = [
("A_no_ema", None), # EMAなし(通常の学習)
("B_ema_0999", 0.999), # EMAあり:ema_momentum=0.999(有効ウィンドウ約1.6epoch相当)
("C_ema_09999", 0.9999), # EMAあり:ema_momentum=0.9999(有効ウィンドウ約16epoch相当)
]
histories, times, scores, ema_scores = {}, {}, {}, {}
for name, momentum in configs:
print(f"\n=== {name} ===")
set_seed(42)
model = build_model(name)
h, t, optimizer, use_ema = compile_and_fit(model, ema_momentum=momentum)
# 通常(EMA未反映のはず)の重みでの評価
final_score = model.evaluate(x_test, y_test, verbose=0)
label = name.split('_', 1)[1]
histories[label] = h
times[label] = t
scores[label] = final_score
print(f"[通常の重み ] test_accuracy:{final_score[1]:.4f}")
# ※ 以下のfinalize_variable_values()は本来「学習後にEMA平均を反映する」ための呼び出しだが、
# fit()を使った場合は最終epoch後に自動でEMA平均が適用済みのため、実際には無意味な操作になる
# (詳細は本文のハマりポイントを参照)。この時点でmodel.get_weights()はすでにEMA平均後の値。
if use_ema:
raw_weights = [w.copy() for w in model.get_weights()] # 実際にはすでにEMA平均後の重み
optimizer.finalize_variable_values(model.trainable_variables)
ema_score = model.evaluate(x_test, y_test, verbose=0)
ema_scores[label] = ema_score
print(f"[EMA重み ] test_accuracy:{ema_score[1]:.4f}(ema_momentum={momentum})")
model.set_weights(raw_weights) # 念のため元に戻しておく
実行結果をクリックして内容を開く
=== A_no_ema === Epoch 1/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 10s 9ms/step - accuracy: 0.2566 - loss: 1.9351 - val_accuracy: 0.3574 - val_loss: 1.7225 Epoch 2/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.3526 - loss: 1.7015 - val_accuracy: 0.4209 - val_loss: 1.5906 Epoch 3/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4224 - loss: 1.5718 - val_accuracy: 0.4744 - val_loss: 1.4647 Epoch 4/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.4655 - loss: 1.4652 - val_accuracy: 0.4958 - val_loss: 1.3931 Epoch 5/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4923 - loss: 1.3909 - val_accuracy: 0.5112 - val_loss: 1.3389 Epoch 6/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5089 - loss: 1.3433 - val_accuracy: 0.5274 - val_loss: 1.2966 Epoch 7/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5252 - loss: 1.3031 - val_accuracy: 0.5419 - val_loss: 1.2581 Epoch 8/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5377 - loss: 1.2701 - val_accuracy: 0.5447 - val_loss: 1.2361 Epoch 9/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5503 - loss: 1.2401 - val_accuracy: 0.5521 - val_loss: 1.2121 Epoch 10/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5582 - loss: 1.2140 - val_accuracy: 0.5651 - val_loss: 1.1989 Epoch 11/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5693 - loss: 1.1847 - val_accuracy: 0.5735 - val_loss: 1.1571 Epoch 12/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5784 - loss: 1.1620 - val_accuracy: 0.5880 - val_loss: 1.1430 Epoch 13/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5867 - loss: 1.1399 - val_accuracy: 0.5904 - val_loss: 1.1315 Epoch 14/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5972 - loss: 1.1186 - val_accuracy: 0.6048 - val_loss: 1.0966 Epoch 15/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6051 - loss: 1.0973 - val_accuracy: 0.6116 - val_loss: 1.0899 Epoch 16/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6119 - loss: 1.0819 - val_accuracy: 0.6164 - val_loss: 1.0669 Epoch 17/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6161 - loss: 1.0622 - val_accuracy: 0.6211 - val_loss: 1.0719 Epoch 18/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6249 - loss: 1.0456 - val_accuracy: 0.6244 - val_loss: 1.0578 Epoch 19/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.6312 - loss: 1.0309 - val_accuracy: 0.6358 - val_loss: 1.0254 Epoch 20/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6367 - loss: 1.0111 - val_accuracy: 0.6423 - val_loss: 1.0189 Epoch 21/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6419 - loss: 1.0016 - val_accuracy: 0.6386 - val_loss: 1.0136 Epoch 22/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6453 - loss: 0.9875 - val_accuracy: 0.6497 - val_loss: 0.9873 Epoch 23/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6485 - loss: 0.9754 - val_accuracy: 0.6555 - val_loss: 0.9696 Epoch 24/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6563 - loss: 0.9630 - val_accuracy: 0.6568 - val_loss: 0.9710 Epoch 25/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6588 - loss: 0.9492 - val_accuracy: 0.6553 - val_loss: 0.9725 Epoch 26/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6651 - loss: 0.9379 - val_accuracy: 0.6584 - val_loss: 0.9738 Epoch 27/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6684 - loss: 0.9252 - val_accuracy: 0.6625 - val_loss: 0.9498 Epoch 28/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6745 - loss: 0.9140 - val_accuracy: 0.6653 - val_loss: 0.9453 Epoch 29/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6776 - loss: 0.9012 - val_accuracy: 0.6712 - val_loss: 0.9299 Epoch 30/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6849 - loss: 0.8878 - val_accuracy: 0.6659 - val_loss: 0.9494 [通常の重み ] test_accuracy:0.6578 === B_ema_0999 === Epoch 1/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 8s 9ms/step - accuracy: 0.2579 - loss: 1.9322 - val_accuracy: 0.3536 - val_loss: 1.7223 Epoch 2/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.3580 - loss: 1.6974 - val_accuracy: 0.4204 - val_loss: 1.5907 Epoch 3/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4217 - loss: 1.5733 - val_accuracy: 0.4711 - val_loss: 1.4759 Epoch 4/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.4620 - loss: 1.4738 - val_accuracy: 0.4919 - val_loss: 1.3976 Epoch 5/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4869 - loss: 1.3978 - val_accuracy: 0.5013 - val_loss: 1.3681 Epoch 6/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5058 - loss: 1.3491 - val_accuracy: 0.5197 - val_loss: 1.3077 Epoch 7/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5224 - loss: 1.3068 - val_accuracy: 0.5321 - val_loss: 1.2741 Epoch 8/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5374 - loss: 1.2706 - val_accuracy: 0.5364 - val_loss: 1.2585 Epoch 9/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5507 - loss: 1.2379 - val_accuracy: 0.5470 - val_loss: 1.2292 Epoch 10/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5606 - loss: 1.2102 - val_accuracy: 0.5595 - val_loss: 1.2010 Epoch 11/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5725 - loss: 1.1819 - val_accuracy: 0.5796 - val_loss: 1.1557 Epoch 12/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5802 - loss: 1.1610 - val_accuracy: 0.5802 - val_loss: 1.1540 Epoch 13/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5892 - loss: 1.1396 - val_accuracy: 0.5906 - val_loss: 1.1273 Epoch 14/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5984 - loss: 1.1144 - val_accuracy: 0.6085 - val_loss: 1.0821 Epoch 15/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6033 - loss: 1.0972 - val_accuracy: 0.6099 - val_loss: 1.0809 Epoch 16/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6110 - loss: 1.0782 - val_accuracy: 0.6122 - val_loss: 1.0631 Epoch 17/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6173 - loss: 1.0582 - val_accuracy: 0.6185 - val_loss: 1.0485 Epoch 18/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6274 - loss: 1.0414 - val_accuracy: 0.6313 - val_loss: 1.0210 Epoch 19/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6334 - loss: 1.0223 - val_accuracy: 0.6308 - val_loss: 1.0191 Epoch 20/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6381 - loss: 1.0050 - val_accuracy: 0.6430 - val_loss: 0.9931 Epoch 21/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6424 - loss: 0.9937 - val_accuracy: 0.6448 - val_loss: 0.9851 Epoch 22/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6485 - loss: 0.9803 - val_accuracy: 0.6531 - val_loss: 0.9657 Epoch 23/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6531 - loss: 0.9648 - val_accuracy: 0.6500 - val_loss: 0.9711 Epoch 24/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6606 - loss: 0.9493 - val_accuracy: 0.6561 - val_loss: 0.9481 Epoch 25/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6610 - loss: 0.9378 - val_accuracy: 0.6578 - val_loss: 0.9509 Epoch 26/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6691 - loss: 0.9263 - val_accuracy: 0.6570 - val_loss: 0.9562 Epoch 27/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6705 - loss: 0.9133 - val_accuracy: 0.6651 - val_loss: 0.9289 Epoch 28/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6775 - loss: 0.9005 - val_accuracy: 0.6691 - val_loss: 0.9252 Epoch 29/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6797 - loss: 0.8875 - val_accuracy: 0.6681 - val_loss: 0.9211 Epoch 30/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6848 - loss: 0.8775 - val_accuracy: 0.6669 - val_loss: 0.9278 [通常の重み ] test_accuracy:0.6884 [EMA重み ] test_accuracy:0.6884(ema_momentum=0.999) === C_ema_09999 === Epoch 1/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 9s 8ms/step - accuracy: 0.2582 - loss: 1.9327 - val_accuracy: 0.3574 - val_loss: 1.7261 Epoch 2/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.3636 - loss: 1.6932 - val_accuracy: 0.4214 - val_loss: 1.5828 Epoch 3/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4219 - loss: 1.5684 - val_accuracy: 0.4698 - val_loss: 1.4671 Epoch 4/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4607 - loss: 1.4708 - val_accuracy: 0.4857 - val_loss: 1.4017 Epoch 5/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4886 - loss: 1.3959 - val_accuracy: 0.5094 - val_loss: 1.3570 Epoch 6/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5059 - loss: 1.3499 - val_accuracy: 0.5290 - val_loss: 1.2938 Epoch 7/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5233 - loss: 1.3093 - val_accuracy: 0.5417 - val_loss: 1.2625 Epoch 8/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5379 - loss: 1.2743 - val_accuracy: 0.5439 - val_loss: 1.2351 Epoch 9/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5507 - loss: 1.2423 - val_accuracy: 0.5574 - val_loss: 1.2092 Epoch 10/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5614 - loss: 1.2117 - val_accuracy: 0.5696 - val_loss: 1.1779 Epoch 11/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5719 - loss: 1.1818 - val_accuracy: 0.5840 - val_loss: 1.1408 Epoch 12/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5801 - loss: 1.1576 - val_accuracy: 0.5875 - val_loss: 1.1306 Epoch 13/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5907 - loss: 1.1378 - val_accuracy: 0.5947 - val_loss: 1.1091 Epoch 14/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.5998 - loss: 1.1123 - val_accuracy: 0.6093 - val_loss: 1.0769 Epoch 15/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6057 - loss: 1.0931 - val_accuracy: 0.6087 - val_loss: 1.0822 Epoch 16/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6111 - loss: 1.0762 - val_accuracy: 0.6195 - val_loss: 1.0544 Epoch 17/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 6ms/step - accuracy: 0.6180 - loss: 1.0570 - val_accuracy: 0.6257 - val_loss: 1.0423 Epoch 18/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6250 - loss: 1.0426 - val_accuracy: 0.6328 - val_loss: 1.0202 Epoch 19/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6308 - loss: 1.0248 - val_accuracy: 0.6392 - val_loss: 1.0034 Epoch 20/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6368 - loss: 1.0122 - val_accuracy: 0.6395 - val_loss: 1.0033 Epoch 21/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6429 - loss: 0.9958 - val_accuracy: 0.6488 - val_loss: 0.9810 Epoch 22/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6477 - loss: 0.9824 - val_accuracy: 0.6524 - val_loss: 0.9676 Epoch 23/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6511 - loss: 0.9709 - val_accuracy: 0.6536 - val_loss: 0.9552 Epoch 24/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6578 - loss: 0.9576 - val_accuracy: 0.6602 - val_loss: 0.9377 Epoch 25/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6617 - loss: 0.9443 - val_accuracy: 0.6588 - val_loss: 0.9465 Epoch 26/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6665 - loss: 0.9341 - val_accuracy: 0.6643 - val_loss: 0.9289 Epoch 27/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6703 - loss: 0.9188 - val_accuracy: 0.6615 - val_loss: 0.9280 Epoch 28/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6763 - loss: 0.9085 - val_accuracy: 0.6698 - val_loss: 0.9088 Epoch 29/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6797 - loss: 0.8971 - val_accuracy: 0.6732 - val_loss: 0.9030 Epoch 30/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6864 - loss: 0.8845 - val_accuracy: 0.6682 - val_loss: 0.9119 [通常の重み ] test_accuracy:0.6347 [EMA重み ] test_accuracy:0.6347(ema_momentum=0.9999)
グラフ+サマリー
# ── val_accuracy / val_loss 比較グラフ ───────────────
fig, axes = plt.subplots(1, 2, figsize=(14, 5))
for label, h in histories.items():
axes[0].plot(h.history['val_accuracy'], label=label)
axes[1].plot(h.history['val_loss'], label=label)
axes[0].set_title('val_accuracy の比較(全30エポック)')
axes[1].set_title('val_loss の比較(全30エポック)')
for ax in axes:
ax.set_xlabel('Epoch'); ax.legend(); ax.grid(True, alpha=0.3)
plt.tight_layout()
plt.savefig('ema_comparison.png', dpi=150)
plt.show()
print("\n===== 最終結果サマリー =====")
print(f"{'Pattern':>14} | {'通常 Acc':>9} | {'EMA Acc':>9} | {'差分':>8} | {'Time(s)':>8}")
print("-" * 60)
for label in ['no_ema', 'ema_0999', 'ema_09999']:
final_acc = scores[label][1]
ema_acc = ema_scores.get(label, (None, None))[1]
diff = (ema_acc - final_acc) if ema_acc is not None else None
t = times[label]
ema_str = f"{ema_acc:.4f}" if ema_acc is not None else " ー"
diff_str = f"{diff:+.4f}" if diff is not None else " ー"
print(f"{label:>14} | {final_acc:>9.4f} | {ema_str:>9} | {diff_str:>8} | {t:>8.1f}")
print("-" * 60)
最終結果サマリー
===== 最終結果サマリー =====
Pattern | 通常 Acc | EMA Acc | 差分 | Time(s)
------------------------------------------------------------
no_ema | 0.6578 | ー | ー | 124.6
ema_0999 | 0.6884 | 0.6884 | +0.0000 | 123.9
ema_09999 | 0.6347 | 0.6347 | +0.0000 | 125.9
------------------------------------------------------------
実験結果
精度グラフ
損失グラフ
| パターン | test_accuracy | 学習時間 |
|---|---|---|
| A:EMAなし | 65.78% | 124.6秒 |
| B:EMAあり(ema_momentum=0.999) | 68.84%(+3.06pt) | 123.9秒 |
| C:EMAあり(ema_momentum=0.9999) | 63.47%(−2.31pt) | 125.9秒 |
B・Cのtest_accuracyは、fit()終了時に自動的に適用されたEMA平均重みでの評価です。実は今回のコードでは学習後にoptimizer.finalize_variable_values()を手動で呼び出していましたが、「通常の重み」と「EMA重み」のtest_accuracyが完全に一致するという結果になりました。これは実装のバグではなく、Kerasの仕様に起因する重要な落とし穴です。詳しくは次のハマりポイントで解説します。
fit()を使う場合、finalize_variable_values()を手動で呼ぶ必要はない:Kerasの公式ドキュメントには「組み込みのfit()ループを使う場合、最終エポック終了後に自動でEMA平均が重みに適用されるため、何もする必要はない」と明記されている。今回のコードでは学習後に明示的にfinalize_variable_values()を呼び出していたが、実際にはその時点ですでに重みはEMA平均済みだった。そのため「通常の重み」として退避したつもりのraw_weightsも、実はすでにEMA平均後の重みであり、手動呼び出しは完全に無意味な操作だった。実験結果で「通常の重み」と「EMA重み」のtest_accuracyが寸分違わず一致したのは、このためである。- 同一の学習run内で「EMA適用前の生の重み」を取得したい場合は、学習途中でコールバックを使って退避する必要がある:
on_epoch_endなどでmodel.get_weights()を都度コピーしておけば、fit()の自動ファイナライズが走る前の状態を確保できる。fit()が完了した後からでは、生の重みにはもうアクセスできない。 ema_momentumはステップ(バッチ)単位で適用される:CIFAR-10で40,000件・batch_size=64の場合、1エポックあたり625ステップ。同じema_momentumの値でも、データ量やバッチサイズが変わると実質的な平滑化の強さ(有効ウィンドウ)が変わる点に注意。
考察
① ema_momentumの値によって「初期のランダムな重み」の残存率が大きく変わる
EMAは学習の最初のステップから適用されるため、$\text{decay}^T$($T$は総ステップ数)の分だけ、学習前の初期化直後のランダムな重みの影響が残り続けます。今回の設定(625ステップ/epoch × 30epoch = 18,750ステップ)で計算すると、次のようになります。
| ema_momentum | 有効ウィンドウ(ステップ換算) | 有効ウィンドウ(epoch換算) | 初期重みの残存率(decay^18750) |
|---|---|---|---|
| 0.999 | 約1,000ステップ | 約1.6epoch | 0.0000007%(無視できる) |
| 0.9999 | 約10,000ステップ | 約16epoch | 約15.3%(無視できない) |
今回の結果はこの理論値ときれいに対応しています。ema_momentum=0.999(B)は初期重みの残存がほぼ無視できるレベルで、EMAなし(65.78%)から+3.06pt改善し68.84%に到達しました。一方ema_momentum=0.9999(C)は初期重みが約15.3%も残ってしまい、EMAなしより−2.31pt悪化し63.47%にとどまりました。学習前のランダムな重み(正解率約10%相当)が15%も混ざれば、最終的な精度が下がるのは理にかなっています。
② SWAの記事の結果と比較して——収束前の学習に対する挙動の違い
前回のSWAの記事では、学習後半10〜20epochの重みを均等平均したところ、いずれも精度が下がりました。原因は、平均対象の区間内でval_accuracyがまだ明確に上昇中で、「未成熟な」重みが平均に混ざり込んでしまうことでした。
今回のEMA(ema_momentum=0.999)は、同じく「まだ収束していない学習」に対して適用したにもかかわらず、精度が向上しました。これはSWAとEMAの仕組みの違いによるものと考えられます。ema_momentum=0.999の有効ウィンドウは約1.6epoch分と非常に短く、直近のごく僅かな範囲の重みを指数的に重み付けして平滑化しているにすぎません。これはSWAのように「学習途中の未成熟な重みを幅広く均等平均する」のではなく、「最終盤の重みに残るバッチごとのノイズを軽くならす」効果に近いといえます。一方、有効ウィンドウを広げすぎたC(ema_momentum=0.9999、約16epoch相当)では、SWAと似た「未成熟な重みの混入」に加えて、EMA特有の「初期のランダムな重みの残存」という別の失敗要因も重なり、SWAの20epoch平均(−2.74pt)に匹敵する悪化(−2.31pt)となりました。
まとめると、EMAは有効ウィンドウを十分短く保てば、収束前の学習でも精度向上に寄与しうる一方、ウィンドウを広げすぎるとSWAと同様の問題に加えて初期重みの残存という固有の問題も抱えるという結論になります。
③ 学習時間に差は出たか
A(EMAなし):124.6秒、B(momentum=0.999):123.9秒、C(momentum=0.9999):125.9秒でした。SWAの記事のときと同様、いずれも誤差の範囲内の差でした。EMAはOptimizerの内部でステップごとに軽量な計算(重みの指数移動平均の更新)を行うだけのため、学習時間への影響はほぼ無視できることが確認できました。
実務での推奨
| 状況 | 推奨 |
|---|---|
| KerasでEMAを試したい | use_ema=Trueとema_momentumを指定するだけでOK。fit()を使うならfinalize_variable_values()の手動呼び出しは不要(自動で適用される) |
| 学習前の「生の重み」も比較・保存したい | 学習中にon_epoch_endコールバックでget_weights()を退避しておく。学習後から取得しようとしてもすでにEMA平均済みで手遅れ |
ema_momentumを決める際 | 総ステップ数に対して$\text{decay}^{T}$が十分小さくなる値を選ぶ。目安は有効ウィンドウ($1/(1-\text{decay})$)が総ステップ数の1割未満になる値 |
| 学習がまだ収束していない設定でどちらを使うか | SWAは平均区間内の未成熟な重みを引きずり込みやすく逆効果になりやすい。EMAはema_momentumを小さめにして有効ウィンドウを短く保てば、精度向上に寄与しやすい |
- EMAはSWAと同じ「重みの平均化」系統だが、直近ほど強く重視する指数的な減衰の仕組みを持つ
- KerasのOptimizerには
use_ema/ema_momentumとして公式にサポートされている。fit()を使う場合、finalize_variable_values()の手動呼び出しは不要(自動適用される。今回のコードでの手動呼び出しは実は無意味な操作だった) - 今回の実験では、
ema_momentum=0.999(有効ウィンドウ約1.6epoch)は精度を+3.06pt改善(65.78%→68.84%)した一方、ema_momentum=0.9999(有効ウィンドウ約16epoch)は初期のランダムな重みの残存(約15.3%)により−2.31pt悪化(65.78%→63.47%)した - SWAが収束前の学習では常に逆効果だったのに対し、EMAは有効ウィンドウを適切に(小さめに)設定すれば収束前でも精度向上に寄与しうる。ただしウィンドウを広げすぎると別の失敗要因(初期重みの残存)が現れる
関連記事もあわせてどうぞ:



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