「Mixed Precisionを使うと学習が速くなる」——そう聞いたことはあっても、精度への影響やどんな条件で効果が出るかまで確認したことはありますか?
今回はGoogle ColabのT4 GPUとCIFAR-10を使い、fp32(通常の32bit精度)とmixed_float16(Mixed Precision)の2パターンで、学習速度・GPUメモリ使用量・test_accuracyを比較しました。速度だけでなく「精度が本当に落ちないのか」「どんな条件だと効果が薄いのか」まで踏み込んで検証します。
なお、Mixed Precisionの基本的な使い方は → Keras Mixed Precision(半精度学習)で高速化する方法|GPUで最大2〜3倍高速化 をご覧ください。本記事はCIFAR-10・小規模CNNでの精度と速度のトレードオフを実験で数値化します。
📘 この記事でわかること
- Mixed Precision(fp32 vs mixed_float16)でtest_accuracy・学習時間・GPUメモリはどう変わるか
- CIFAR-10のような小さい入力・小規模CNNでも速度向上の恩恵が出るのか
- 出力層をfloat32に固定しないと何が起きるか(数値不安定の実例)
Mixed Precisionとは何をしているのか
Mixed Precisionは、重みの保持はfloat32のまま、演算の一部をfloat16(半精度)で行うことで、GPUのTensor Coreを活用し計算を高速化する仕組みです。float16はfloat32に比べて表現できる指数部のビット数が少なく、扱える数値の範囲が狭くなります。
$$ \text{float32} = 1\,(\text{符号}) + 8\,(\text{指数部}) + 23\,(\text{仮数部}) \text{ bit} \\ \text{float16} = 1\,(\text{符号}) + 5\,(\text{指数部}) + 10\,(\text{仮数部}) \text{ bit} $$
指数部が8bitから5bitに減ることで、float16が表現できる数値の範囲はfloat32よりも大幅に狭くなります。この範囲の狭さが、後述する「出力層をfloat32にすべき理由」に直結します。
| 項目 | fp32(通常) | mixed_float16 |
|---|---|---|
| 重みの保持形式 | float32 | float32(マスター重み) |
| 順伝播・逆伝播の演算 | float32 | float16(Tensor Core活用) |
| 勾配のスケーリング | 不要 | LossScaleOptimizerが自動で実施 |
| 期待される効果 | 基準 | 学習速度向上・GPUメモリ削減 |
| リスク | なし | 数値のオーバーフロー・アンダーフロー |
T4 GPU(Turing世代、Compute Capability 7.5)はTensor Coreを搭載しているため、理論上はMixed Precisionの恩恵を受けられる環境です。ただし、CIFAR-10は32×32という小さい入力サイズで、かつ今回のCNNも2層と小規模なため、Tensor Coreの恩恵がどこまで出るかは実際に測ってみないとわかりません。この「小規模モデルでの効果の有無」が本記事の検証ポイントです。
実験コード
使用環境はGoogle Colab(GPU:T4)、データセットはCIFAR-10です。Mixed Precisionのポリシー設定以外は全て同一にして、精度・速度への影響だけを取り出します。
環境準備(最初に一度だけ実行)
# ── 環境準備(最初に一度だけ実行)──────────────────────
!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 3s (2,625 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 114.1 MB/s eta 0:00:00
Preparing metadata (setup.py) ... done
Building wheel for japanize_matplotlib (setup.py) ... done
環境準備完了
GPU環境の確認
Mixed Precisionの効果はGPUのCompute Capabilityに依存するため、まず環境を確認します。
import tensorflow as tf
gpus = tf.config.list_physical_devices('GPU')
print("GPU:", gpus)
for gpu in gpus:
details = tf.config.experimental.get_device_details(gpu)
print("Compute Capability:", details.get('compute_capability'))
実行結果をクリックして内容を開く
GPU: [PhysicalDevice(name='/physical_device:GPU:0', device_type='GPU')] Compute Capability: (7, 5)
import・データ準備・モデル構築関数
import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import mixed_precision
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 build_model(policy_name, name):
# ポリシーはグローバル設定のため、モデル構築前に毎回明示的に切り替える
mixed_precision.set_global_policy(policy_name)
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),
# 出力層は数値安定のため明示的にfloat32へ固定
keras.layers.Dense(10, activation='softmax', dtype='float32'),
], name=name)
return model
def compile_and_fit(model):
model.compile(optimizer='adam',
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)
return history, time.time() - start
実行結果をクリックして内容を開く
Downloading data from https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz 170498071/170498071 ━━━━━━━━━━━━━━━━━━━━ 1324s 8us/step
2パターンの学習実行(fp32 vs mixed_float16)
configs = [('float32', 'A_fp32'), ('mixed_float16', 'B_mixed_fp16')]
histories, times, scores, params, peak_mem = {}, {}, {}, {}, {}
for policy_name, name in configs:
print(f"\n=== {name}(policy={policy_name}) ===")
# GPUメモリ計測をリセット
tf.config.experimental.reset_memory_stats('GPU:0')
model = build_model(policy_name, name)
print("compute dtype:", model.dtype_policy.compute_dtype,
" / variable dtype:", model.dtype_policy.variable_dtype)
print(model.summary())
h, t = compile_and_fit(model)
s = model.evaluate(x_test, y_test, verbose=0)
mem_info = tf.config.experimental.get_memory_info('GPU:0')
label = name.split('_', 1)[1]
histories[label] = h
times[label] = t
scores[label] = s
params[label] = model.count_params()
peak_mem[label] = mem_info['peak'] / (1024 ** 2) # MB換算
print(f"学習時間:{t:.1f}秒 test_accuracy:{s[1]:.4f} "
f"GPUピークメモリ:{peak_mem[label]:.1f}MB")
# 次の実験に影響を与えないよう、ポリシーをデフォルトに戻す
mixed_precision.set_global_policy('float32')
実行結果をクリックして内容を開く
=== A_fp32(policy=float32) === compute dtype: float32 / variable dtype: float32 Model: "A_fp32" ┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━┓ ┃ Layer (type) ┃ Output Shape ┃ Param # ┃ ┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━┩ │ conv2d (Conv2D) │ (None, 32, 32, 64) │ 1,792 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ max_pooling2d (MaxPooling2D) │ (None, 16, 16, 64) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ conv2d_1 (Conv2D) │ (None, 16, 16, 128) │ 73,856 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ max_pooling2d_1 (MaxPooling2D) │ (None, 8, 8, 128) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ global_average_pooling2d │ (None, 128) │ 0 │ │ (GlobalAveragePooling2D) │ │ │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dense (Dense) │ (None, 128) │ 16,512 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dropout (Dropout) │ (None, 128) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dense_1 (Dense) │ (None, 10) │ 1,290 │ └─────────────────────────────────┴────────────────────────┴───────────────┘ Total params: 93,450 (365.04 KB) Trainable params: 93,450 (365.04 KB) Non-trainable params: 0 (0.00 B) None Epoch 1/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 10s 9ms/step - accuracy: 0.2612 - loss: 1.9263 - val_accuracy: 0.3526 - val_loss: 1.7196 Epoch 2/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.3674 - loss: 1.6850 - val_accuracy: 0.4257 - val_loss: 1.5713 Epoch 3/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4229 - loss: 1.5710 - val_accuracy: 0.4477 - val_loss: 1.5220 Epoch 4/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4648 - loss: 1.4662 - val_accuracy: 0.4832 - val_loss: 1.4215 Epoch 5/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.4942 - loss: 1.3951 - val_accuracy: 0.5112 - val_loss: 1.3394 Epoch 6/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5087 - loss: 1.3470 - val_accuracy: 0.5199 - val_loss: 1.3076 Epoch 7/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5267 - loss: 1.2966 - val_accuracy: 0.5264 - val_loss: 1.2889 Epoch 8/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5389 - loss: 1.2666 - val_accuracy: 0.5481 - val_loss: 1.2505 Epoch 9/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5519 - loss: 1.2327 - val_accuracy: 0.5558 - val_loss: 1.2005 Epoch 10/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5617 - loss: 1.2076 - val_accuracy: 0.5568 - val_loss: 1.2114 Epoch 11/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5747 - loss: 1.1784 - val_accuracy: 0.5747 - val_loss: 1.1664 Epoch 12/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5793 - loss: 1.1610 - val_accuracy: 0.5848 - val_loss: 1.1422 Epoch 13/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.5925 - loss: 1.1355 - val_accuracy: 0.5969 - val_loss: 1.1022 Epoch 14/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.5975 - loss: 1.1148 - val_accuracy: 0.6043 - val_loss: 1.0996 Epoch 15/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6087 - loss: 1.0890 - val_accuracy: 0.6081 - val_loss: 1.0803 Epoch 16/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6094 - loss: 1.0786 - val_accuracy: 0.6218 - val_loss: 1.0412 Epoch 17/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6190 - loss: 1.0588 - val_accuracy: 0.6143 - val_loss: 1.0515 Epoch 18/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6249 - loss: 1.0406 - val_accuracy: 0.6320 - val_loss: 1.0322 Epoch 19/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6309 - loss: 1.0298 - val_accuracy: 0.6277 - val_loss: 1.0304 Epoch 20/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6400 - loss: 1.0107 - val_accuracy: 0.6210 - val_loss: 1.0298 Epoch 21/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6422 - loss: 0.9941 - val_accuracy: 0.6442 - val_loss: 0.9851 Epoch 22/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6460 - loss: 0.9831 - val_accuracy: 0.6515 - val_loss: 0.9693 Epoch 23/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6550 - loss: 0.9640 - val_accuracy: 0.6555 - val_loss: 0.9488 Epoch 24/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6571 - loss: 0.9567 - val_accuracy: 0.6461 - val_loss: 0.9725 Epoch 25/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6646 - loss: 0.9355 - val_accuracy: 0.6586 - val_loss: 0.9550 Epoch 26/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6647 - loss: 0.9250 - val_accuracy: 0.6481 - val_loss: 0.9602 Epoch 27/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6726 - loss: 0.9156 - val_accuracy: 0.6699 - val_loss: 0.9196 Epoch 28/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6754 - loss: 0.9033 - val_accuracy: 0.6596 - val_loss: 0.9532 Epoch 29/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6802 - loss: 0.8912 - val_accuracy: 0.6710 - val_loss: 0.9114 Epoch 30/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 7ms/step - accuracy: 0.6831 - loss: 0.8814 - val_accuracy: 0.6815 - val_loss: 0.8938 学習時間:122.6秒 test_accuracy:0.6801 GPUピークメモリ:1009.5MB === B_mixed_fp16(policy=mixed_float16) === compute dtype: float16 / variable dtype: float32 Model: "B_mixed_fp16" ┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━┓ ┃ Layer (type) ┃ Output Shape ┃ Param # ┃ ┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━┩ │ conv2d_2 (Conv2D) │ (None, 32, 32, 64) │ 1,792 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ max_pooling2d_2 (MaxPooling2D) │ (None, 16, 16, 64) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ conv2d_3 (Conv2D) │ (None, 16, 16, 128) │ 73,856 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ max_pooling2d_3 (MaxPooling2D) │ (None, 8, 8, 128) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ global_average_pooling2d_1 │ (None, 128) │ 0 │ │ (GlobalAveragePooling2D) │ │ │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dense_2 (Dense) │ (None, 128) │ 16,512 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dropout_1 (Dropout) │ (None, 128) │ 0 │ ├─────────────────────────────────┼────────────────────────┼───────────────┤ │ dense_3 (Dense) │ (None, 10) │ 1,290 │ └─────────────────────────────────┴────────────────────────┴───────────────┘ Total params: 93,450 (365.04 KB) Trainable params: 93,450 (365.04 KB) Non-trainable params: 0 (0.00 B) None Epoch 1/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 13s 8ms/step - accuracy: 0.2663 - loss: 1.9165 - val_accuracy: 0.3368 - val_loss: 1.7930 Epoch 2/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.3709 - loss: 1.6712 - val_accuracy: 0.4045 - val_loss: 1.6011 Epoch 3/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.4284 - loss: 1.5413 - val_accuracy: 0.4067 - val_loss: 1.6143 Epoch 4/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 6ms/step - accuracy: 0.4654 - loss: 1.4528 - val_accuracy: 0.4753 - val_loss: 1.4474 Epoch 5/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.4902 - loss: 1.3901 - val_accuracy: 0.5193 - val_loss: 1.3222 Epoch 6/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5112 - loss: 1.3413 - val_accuracy: 0.5136 - val_loss: 1.3192 Epoch 7/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5255 - loss: 1.3021 - val_accuracy: 0.5396 - val_loss: 1.2627 Epoch 8/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 5ms/step - accuracy: 0.5393 - loss: 1.2670 - val_accuracy: 0.5366 - val_loss: 1.2631 Epoch 9/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5508 - loss: 1.2385 - val_accuracy: 0.5515 - val_loss: 1.2085 Epoch 10/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5606 - loss: 1.2134 - val_accuracy: 0.5686 - val_loss: 1.1709 Epoch 11/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 6ms/step - accuracy: 0.5692 - loss: 1.1880 - val_accuracy: 0.5821 - val_loss: 1.1447 Epoch 12/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5774 - loss: 1.1668 - val_accuracy: 0.5904 - val_loss: 1.1308 Epoch 13/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5867 - loss: 1.1424 - val_accuracy: 0.5933 - val_loss: 1.1175 Epoch 14/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.5951 - loss: 1.1139 - val_accuracy: 0.5917 - val_loss: 1.1102 Epoch 15/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 4s 6ms/step - accuracy: 0.6016 - loss: 1.0984 - val_accuracy: 0.5865 - val_loss: 1.1133 Epoch 16/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 5s 5ms/step - accuracy: 0.6075 - loss: 1.0812 - val_accuracy: 0.6077 - val_loss: 1.0652 Epoch 17/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6201 - loss: 1.0552 - val_accuracy: 0.6199 - val_loss: 1.0337 Epoch 18/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6257 - loss: 1.0423 - val_accuracy: 0.6079 - val_loss: 1.0696 Epoch 19/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6300 - loss: 1.0260 - val_accuracy: 0.6373 - val_loss: 0.9998 Epoch 20/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6374 - loss: 1.0089 - val_accuracy: 0.6366 - val_loss: 0.9937 Epoch 21/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6410 - loss: 0.9953 - val_accuracy: 0.6423 - val_loss: 0.9826 Epoch 22/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6466 - loss: 0.9781 - val_accuracy: 0.6461 - val_loss: 0.9854 Epoch 23/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6557 - loss: 0.9617 - val_accuracy: 0.6540 - val_loss: 0.9536 Epoch 24/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6564 - loss: 0.9481 - val_accuracy: 0.6433 - val_loss: 0.9942 Epoch 25/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6653 - loss: 0.9366 - val_accuracy: 0.6553 - val_loss: 0.9564 Epoch 26/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6658 - loss: 0.9288 - val_accuracy: 0.6669 - val_loss: 0.9279 Epoch 27/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6763 - loss: 0.9085 - val_accuracy: 0.6684 - val_loss: 0.9246 Epoch 28/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6763 - loss: 0.8969 - val_accuracy: 0.6745 - val_loss: 0.9132 Epoch 29/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6842 - loss: 0.8891 - val_accuracy: 0.6664 - val_loss: 0.9163 Epoch 30/30 625/625 ━━━━━━━━━━━━━━━━━━━━ 3s 5ms/step - accuracy: 0.6895 - loss: 0.8754 - val_accuracy: 0.6792 - val_loss: 0.9006 学習時間:106.3秒 test_accuracy:0.6776 GPUピークメモリ:787.7MB
グラフ+サマリー
# ── 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('mixed_precision_comparison.png', dpi=150)
plt.show()
print("\n===== 最終結果サマリー =====")
print(f"{'Pattern':>12} | {'Val Acc':>8} | {'Test Acc':>9} | {'Time(s)':>8} | {'Peak Mem(MB)':>13}")
print("-" * 62)
for label in ['fp32', 'mixed_fp16']:
val_acc = histories[label].history['val_accuracy'][-1]
test_acc = scores[label][1]
t = times[label]
m = peak_mem[label]
print(f"{label:>12} | {val_acc:>8.4f} | {test_acc:>9.4f} | {t:>8.1f} | {m:>13.1f}")
print("-" * 62)
実行結果をクリックして内容を開く
===== 最終結果サマリー =====
Pattern | Val Acc | Test Acc | Time(s) | Peak Mem(MB)
--------------------------------------------------------------
fp32 | 0.6815 | 0.6801 | 122.6 | 1009.5
mixed_fp16 | 0.6792 | 0.6776 | 106.3 | 787.7
--------------------------------------------------------------
実験結果
精度グラフ
損失グラフ
結果サマリー
| パターン | 最終 val_accuracy | 最終 test_accuracy | 学習時間 | GPUピークメモリ |
|---|---|---|---|---|
| A:fp32(通常) | 68.15% | 68.01% | 122.6秒 | 1009.5MB |
| B:mixed_float16 | 67.92% | 67.76% | 106.3秒 | 787.7MB |
学習時間は13.3%短縮(122.6秒→106.3秒)、GPUピークメモリは22.0%削減(1009.5MB→787.7MB)、test_accuracyの差は-0.25pt(68.01%→67.76%)に留まりました。
ハマりポイント
- 出力層をfloat32に固定しないとsoftmaxの計算が不安定になる。mixed_float16のままsoftmax出力を計算すると、float16の狭い表現範囲でオーバーフロー・アンダーフローが起き、loss=NaNになることがある。
Dense(10, activation='softmax', dtype='float32')のように出力層だけ明示的にfloat32へ固定するのが定石。 - mixed_precision.set_global_policy()はグローバル設定。2パターン目のモデルを作る前に元のポリシーへ戻し忘れると、fp32のつもりのモデルが実はmixed_float16のままになる。ループの最後で必ず
set_global_policy('float32')に戻すこと。 - BatchNormalizationはfloat32のまま扱われる。Keras側が数値安定性のためBatchNorm層の計算を自動的にfloat32で行うため、Mixed Precisionを使ってもBatchNorm自体の挙動は変わらない(今回はBatchNorm未使用構成だが、追加する場合は覚えておきたい)。
- 小規模モデル・小さい入力サイズでは速度向上が限定的な場合がある。Tensor Coreの恩恵は行列演算のサイズが大きいほど出やすい。CIFAR-10(32×32)・Conv2D 2層という小規模構成では、データ転送やPythonループのオーバーヘッドが相対的に大きく、期待したほどの高速化が出ない可能性がある。
考察
① 学習時間はどの程度短縮されたか
学習時間はfp32の122.6秒に対しmixed_float16は106.3秒で、13.3%の短縮にとどまりました。Mixed Precisionは「2〜3倍高速化」と紹介されることが多いですが、それはResNetやTransformerのような大きな行列演算を大量に含むモデルでの数値です。今回のCNNはConv2D×2層・Dense(128)というごく小規模な構成で、かつ入力も32×32と小さいため、1回あたりの演算がTensor Coreの並列処理能力を使い切れず、恩恵が限定的になったと考えられます。バッチサイズ64・小規模モデルという条件では、Mixed Precisionの速度メリットは「劇的」ではなく「そこそこ」と捉えるのが実態に近いと言えます。
② test_accuracyへの影響
test_accuracyはfp32が68.01%、mixed_float16が67.76%で、差は-0.25ptでした。この程度の差はシード未固定の実行間ばらつきの範囲内に収まる水準であり、「Mixed Precisionにしたから精度が落ちた」と結論づけるには根拠不足です。出力層をdtype='float32'で固定したことで、softmaxの数値不安定によるNaNや大きな精度劣化は発生しませんでした。この結果から、出力層をfloat32に固定する対策は今回のCIFAR-10実験で有効に機能したと言えます。
③ GPUメモリ使用量への影響
GPUピークメモリはfp32が1009.5MB、mixed_float16が787.7MBで、22.0%の削減となりました。速度向上(13.3%)よりもメモリ削減幅(22.0%)の方が大きく出ている点が今回の実験の特徴です。これはfloat16化によってアクティベーション(中間層の出力)のメモリ使用量が半分近くまで縮小する一方、演算速度そのものは小規模モデルゆえにTensor Coreの並列度をフルに引き出せなかったためと考えられます。「速度より先にメモリ面でのメリットが出る」のは、CIFAR-10規模の小さいモデルでMixed Precisionを使う際の実務的な着目点と言えます。
④ 予想外の結果が出た場合の確認ポイント
今回の実験では、速度・メモリともに期待通りの方向(mixed_float16が有利)に動き、精度差も誤差範囲内に収まったため、以下のような異常は発生しませんでした。ただし、環境やモデル構成によっては次のようなケースが起こり得るため、もしmixed_float16の方が明確に遅い、または精度が大きく劣化していた場合は、以下を優先して確認してください。
- Colab環境のGPUが実際にT4(Compute Capability 7.0以上)になっているか(無償枠ではGPUが割り当てられないこともある)
- 出力層に
dtype='float32'が正しく設定されているか model.dtype_policy.compute_dtypeの出力が意図した精度になっているか- 学習曲線にNaN・異常なスパイクが出ていないか
| ケース | 実務での推奨 |
|---|---|
| Tensor Core対応GPU(T4以上)を使える | Mixed Precisionを試す価値あり。出力層のfloat32固定を忘れずに。 |
| モデル・バッチサイズが小さい | 速度向上が限定的な場合がある。バッチサイズを大きくしてから再検証する。 |
| GPUメモリが厳しい | 速度よりメモリ削減目的でMixed Precisionを検討する価値がある。 |
| CPU環境・古いGPU | Tensor Coreの恩恵がなく、むしろ遅くなる場合があるため非推奨。 |
まとめ
CIFAR-10・Conv2D×2層という小規模なCNNでは、Mixed Precisionによる学習時間の短縮は13.3%と控えめでしたが、GPUピークメモリは22.0%削減され、test_accuracyの差は-0.25ptと誤差範囲に収まりました。「劇的な高速化」を狙うより、「精度をほぼ落とさずにメモリを浮かせる」目的で使う方が、小規模モデルでは実態に合っていると言えそうです。より大きなモデルやバッチサイズで実験すれば、速度面のメリットもさらに顕著になる可能性があります。出力層をdtype='float32'に固定する対策は、今回の規模でも有効に機能しました。
関連記事もあわせてどうぞ:
- Mixed Precisionの基本 → Keras Mixed Precision(半精度学習)で高速化する方法|GPUで最大2〜3倍高速化
- バッチサイズの影響 → バッチサイズを変えると精度はどう変わる?(16 vs 64 vs 256)【Keras×CIFAR-10実験】
- GPU・環境確認 → 【2026年版】Google ColabでGPU・TPUを使う方法|設定・確認・トラブル対処まとめ
- Dense層ユニット数比較 → Dense層のユニット数を変えると精度はどう変わる?(32 vs 128 vs 512)【Keras×CIFAR-10実験】



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