深層学習におけるトレーニング高速化とメモリ最適化技術

異なる精度でのトレーニング手法

単精度トレーニング(FP32)は全てのパラメータ、活性値、勾配を32ビット浮動小数点数で表現します。
半精度トレーニング(FP16/BF16)は16ビット浮動小数点数を使用し、FP16はIEEE標準、BF16はAI計算向けに設計された形式です。
混合精度トレーニングはFP16/BF16とFP32を組み合わせ、通常は重みと勾配をFP32、活性値と中間計算をFP16/BF16で処理します。

精度比較表

指標 単精度(FP32) 半精度(FP16/BF16) 混合精度
数値精度 低(FP16)、中(BF16) 中高
メモリ使用量
処理速度 低速 高速 高速
安定性 最良 低(FP16) 良好

混合精度の必要性

FP16単独使用では数値表現範囲の制約からオーバーフロー/アンダーフローが発生しやすくなります。FP16の有効桁数(約10ビット)が少ないため、微小な勾配がゼロに丸められ、学習が停滞する可能性があります。

混合精度解決策

  • FP32マスターコピー:重みの正確なバージョンをFP32で維持し、FP16コピーで高速計算
  • 損失スケーリング:勾配消失問題対策として損失値を拡大・縮小する動的スケーリング

実装例

Apexによる混合精度

# Apexインストール
!git clone https://github.com/NVIDIA/apex
%cd apex
!pip install -v --no-cache-dir .

# トレーニング実装
from apex import amp

network, optim = amp.initialize(network, optim, opt_level="O1")
...
with amp.scale_loss(loss_val, optim) as adjusted_loss:
    adjusted_loss.backward()

PyTorch AMP実装

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()
...
with autocast():
    predictions = network(inputs)
    loss_val = criterion(predictions, labels)
    
scaler.scale(loss_val).backward()
scaler.step(optim)
scaler.update()

性能比較結果

CIFAR10データセット(ResNet50)
手法 精度(訓練) 精度(テスト) 所要時間 メモリ使用量(MB)
PyTorch AMP 0.9364 0.8031 16.99分 15508
Apex 0.9366 0.7956 16.51分 13166
FP32 0.9456 0.8092 22.27分 22818

勾配チェックポイントによるメモリ最適化

この手法は中間活性値の保存領域を削減するため、前方伝播の特定ポイントのみを保存し、勾配計算時に必要な部分を再計算します。メモリ使用量を最大75%削減可能ですが、再計算による処理時間増加とのトレードオフが存在します。

タグ: 混合精度トレーニング 勾配チェックポイント 深層学習 Apex PyTorch

8月2日 21:28 投稿