⏱️

Tesla T4でMNIST推論 2,780万枚/秒を出すための最適化技術

に公開6

本稿は、提示されたColabノートブック(mnist_t4_ultrafast_inference_v7.ipynb)の内容をベースに、「6年前世代のGPU(Tesla T4)」でMNIST推論を毎秒2,800万枚級で回すために効いた最適化を、再現可能な形で体系立てて整理するためのものです。

https://colab.research.google.com/drive/1JiZOW13kX9RYjwhmxesMHzDB3VPFtknq?usp=sharing


1. 実測結果(事実)

このノートブックはGoogle ColabのT4ランタイムで実行するためのものです。
実行環境は下のようなものです。

実行環境

  • GPU: Tesla T4

  • PyTorch: 2.9.0+cu126

  • CUDA: 12.6

  • ベンチ入力:

    • Throughput: x_thr.shape = (1024, 784)(batch=1024)
    • Latency: x_lat.shape = (1, 784)(batch=1)
  • 入力dtype: float16(GPU常駐)

ベースライン vs 最適化済み(事実)

モデル 精度(Accuracy) スループット(wall) レイテンシ(wall)
ベースライン(Ridge-ELM, hidden=512) 0.9239 27,818,225.28 samples/s 10.88 µs/iter
最適化済み(全層CE微調整 + compile選択) 0.9679 24,641,180.85 samples/s 9.22 µs/iter
  • 「2,800万枚/秒級」はベースラインが達成しています(約2.78e7 samples/s)。
  • このバージョンの主眼は「精度を上げつつ速度低下を最小化」で、Accuracyを 0.9239 → 0.9679 へ引き上げています(その代償としてthroughputは約11%低下)。

2. まず前提:この速度は「何を速くした」結果なのか

MNIST推論で2,800万枚/秒級を出すには、単にGPUを使うだけでは足りません。ノートブックの設計思想は以下のようなものがあります。

  • 推論計算をほぼGEMM(行列積)2発 + ReLUに寄せる
  • データ転送・Pythonループ・動的shapeなどの“周辺オーバーヘッド”を削り切る

この2点に集約されます。

スループットが伸びる構造(推論)

MNIST 28×28 を flat 784 にし、

  • h = ReLU( x @ W1^T + b1 )(784→H)
  • y = h @ W2^T + b2(H→C)

2層MLPで完結させます(Conv無し)。これにより、推論はほぼ **cuBLASのGEMM(FP16)**に落ちます。


3. コア戦略:推論パスを「Tensor Coreが走りやすいGEMM」に固定する

3.1 FP16化(実装 + 理由)

実装

  • 入力 x_train/x_testfloat16でGPU常駐させています。
  • 重みも float16を基本にしています(CPU保存もFP16)。

理由

  • Tesla T4はTensor Coreを持つ(Volta/Turing世代)ため、FP16 GEMMで実効TFLOPSが跳ねる
  • MNISTのような小規模モデルは、FP32で回すと計算資源を使い切る前にオーバーヘッドが支配的になりやすい。

再現ポイント

  • 入力のFP16化は「毎step」ではなく「1回だけ」行う(後述のGPU常駐前処理)。

3.2 shape・レイアウトを固定する

実装

  • x_test[:1024].contiguous()x_thr として切り出し、固定shapeの入力を作っています。
  • さらに x_lat = x_test[:1].contiguous()batch=1の固定shapeも用意。
  • 重みも .contiguous() を多用し、転置を使う場合は事前に W1T/W2T を作ってbuffer化しています。

理由

  • torch.compile(fullgraph=True, dynamic=False) や CUDA Graphは、入力shapeが固定されているほど安定し、不要なrecompile/capture失敗を避けることが狙いです。
  • contiguousでないと内部で暗黙copyが入ったり、GEMMが非効率な経路に落ちることがある。

4. オーバーヘッドを殺す:CUDA Graphが「2,800万枚/秒級」の決定打

ここが最重要です。

4.1 microbench(通常実行)では 1,012万 samples/s 程度(事実)

Auto-tune中(microbench)で hidden=512 の候補が示すthroughputは:

  • 10,126,842.66 samples/s(wall)(hidden=512, best pad=10, flinear)

4.2 最終ベンチ(CUDA Graph replay)で 2,781万 samples/s(事実)

最終結果のベースラインthroughputは:

  • 27,818,225.28 samples/s(wall)
    になっています。

4.3 何が起きているか

事実

  • ノートブックは CUDAGraphRunner を作り、固定形状推論を capture → replay で回しています。
  • ベンチ関数はCUDA eventでGPU時間を測りつつ、wall timeも測っています。

理由

  • このモデルは「GEMM2回 + ReLU」という短いカーネル列で、しかもバッチを固定して何度も繰り返します。
  • こういうケースでは、Pythonから毎回カーネルを投入するコスト(launch overhead、dispatcher、フレームワークの管理)が相対的に効いてきます。
  • CUDA Graph replayは、その“毎回の投入”をまとめて録画し、replayを軽量にするため、wall throughputが跳ねます。

実際、microbench→最終で 約2.75倍に上がっています:

  • 27,818,225 / 10,126,842 ≒ 2.75x

2,800万枚/秒級に到達させた最大要因は、GPU計算そのものよりも **「CPU側の発行オーバーヘッド除去(CUDA Graph)」**です。


5. モデル面の工夫:ELM(Ridge)で「探索を速く」、CE微調整で「精度を上げる」

速度最適化だけなら学習は不要ですが、ノートブックは精度も追う設計です。ポイントは、推論コストを増やさずに精度を上げるための学習戦略です。

5.1 Ridge-ELM(Extreme Learning Machine)で初期解を高速に作る(事実)

事実(ノートブック)

  • 1層目 (W1, b1) をランダム固定し、2層目 (W2, b2) を ridge regression で解いています。
  • Gram行列 S = h^T h などの統計量を fp32で蓄積(オーバーフロー/誤差対策)。
  • バイアス b2 は「入力に定数1を追加した回帰」と等価ですが、巨大なconcatを避け、sum_h / sum_y / nでブロック行列を組んで解く実装になっています。
  • solveは cholesky_ex + jitter で頑健化し、ダメなら solve にfallback。

狙い

  • hidden次元やpad、実装差(F.linear vs mm)を探索するには、学習が重いと回りません。
  • Ridge-ELMは「訓練コストを抑えつつ、それなりの精度」を得るのに向き、探索ループ(Auto-tune)を成立させる役割を持ちます。

5.2 Biasを入れて精度を稼ぐ

事実

  • USE_BIAS=True(隠れ層・出力層ともbiasあり)。
  • 推論ではbias加算が入るだけ。

理由

  • MNIST程度のタスクでも、bias無しは表現力が落ちやすい。
  • GEMM支配のパスでは、bias addは相対コストが小さく、精度の上がり幅に対して“コスパが良い”。

5.3 このバージョンの主改善:全層CE微調整

ログの結果

  • Ridge直後の精度: 0.9239
  • 全層CE微調整後の精度: 0.9679
  • 推論の形(2 GEMM + ReLU)は同じで、推論計算量はほぼ不変。

理由

  • ELMの限界は「1層目がランダム固定」な点にあり、分類境界が最適化されません。
  • そこで 推論構造を変えず、学習だけ増やして (W1,b1,W2,b2) をCEで最適化し、精度を引き上げています。
  • これは「推論は速く、精度も高い」を両立する典型手です(推論計算量が増えないため)。

実装上の注意

  • Fine-tune用のパラメータはfloat32で保持し、AMP/autocast下では入力dtypeに合わせて .to(x_flat.dtype) を明示する作りになっています(安定性とGradScalerの都合)。

6. 速度のための実装ディテール:F.linear vs mm、pad=10 vs 16、inplace vs oop ReLU

ノートブックは「実環境依存」を前提に、複数候補を実測で選びます。

6.1 F.linear vs mm

Auto-tune結果

  • 最速候補のvariantは、今回の環境では全て flinear が選ばれています。

理由

  • F.linear は内部で最適なGEMM経路を選びやすく、bias融合(addmm)も絡む。
  • 一方、mm 版は転置済み重みで余計なtransposeフラグを避ける狙いですが、今回の環境では勝てなかった。

6.2 出力クラス次元のパディング(10 vs 16)

  • PAD_CLASSES_CANDIDATES = [10, 16] を試し、速度で選ぶ。

事実(Auto-tune出力)

  • hidden=768/1024/1536/2048/4096 では best pad=16 が選ばれている。
  • hidden=512/3072/6144 では best pad=10 が選ばれている。

理由

  • Tensor CoreやGEMMカーネル選択は、M/N/Kの割り切れ(特に8や16)に強く依存します。
  • C=10 は半端なので C=16 に0埋めして計算し、最後に y[:, :10] で戻す戦略は合理的。(だと思った)
  • ただし “sliceのコスト” と “GEMMの速さ” のトレードがあり、hidden次元や実装により勝敗が変わります。

6.3 in-place ReLU と out-of-place ReLU

ノートブック実装

  • eager推論は torch.relu_(in-place)も選べる。
  • compile探索時には TRY_OOP_RELU_FOR_COMPILE=True として out-of-place ReLU 版も候補に入れる。
  • compile選択の結果、throughput最良は oop + max-autotune-no-cudagraphs

理由

  • eagerではin-placeが有利なことが多い(メモリアロケーションが減る)。
  • しかしコンパイラ最適化では、in-placeがfusionやスケジューリングの邪魔になる場合があり、oopのほうが速いケースが出ます。
  • そのため、**「eagerに最適」≠「compileに最適」**を分けて検証しているのがポイントです。

7. torch.compile:固定形状で“実測選択”する(そしてハマりどころも押さえる)

7.1 コンパイル戦略(事実)

事実

  • torch.compilefullgraph=True(可能なら)かつ dynamic=False(可能なら)で試します。

  • modeは以下を探索:

    • default
    • reduce-overhead
    • max-autotune
    • max-autotune-no-cudagraphs

事実(実測選択)

  • best throughput: oop + max-autotune-no-cudagraphs(≈24.76M samples/s)
  • best latency: inplace + default(≈9.65 µs/iter)

7.2 重要な落とし穴:shape混在でrecompile地獄

ログに size mismatch expected 1 actual 1024 のようなrecompile警告が出ています。ノートブック内でもコメントで強調されている通り、

  • 同じcompiled wrapperに別shapeを流すとrecompileが走り得る

対策としてノートブックは、

  • throughput用(batch=1024)と latency用(batch=1)で 別々にcompileしています。

これをやらないと、パフォーマンスは落ちるだけでなく、recompile_limitに到達して失敗することもあります(実際、ベースライン側のcompileが最終ベンチで失敗しています)。


8. Accuracy-aware Auto-tune:速度ガードレール付きで“使える最適解”を選ぶ

このノートブックの設計が実務的なのは、単に最速を狙うのではなく、

  • 速度を落としすぎる候補を自動的に棄却しつつ
  • 精度目標に応じて最終構成を決める

という選び方をしている点です。

8.1 ガードレール(事実)

  • SPEED_RATIO_MIN=0.80
  • baseline(hidden=512)のthroughputを基準に、80%以上の候補だけ残す。

今回のAuto-tuneログでは、

  • baseline microbench thr ≈ 10.13M
  • floor ≈ 8.10M
    となり、hiddenを増やすと一気にfloor割れして、最終的に hidden=512が選ばれる結果になっています。

8.2 精度の上げ方を「学習側」に寄せる

  • hiddenを増やして精度を稼ぐ代わりに、v7では 全層CE微調整で精度を稼ぐ。
  • 推論の計算グラフ(2 GEMM + ReLU)は同一で、推論を重くしない。

9. 再現のためのポイント

9.1 速度2,800万枚/秒級に効く順

ほぼ確実に効く(必須級)

  1. 入力をGPU常駐にする(転送ゼロ)
  2. FP16(入力・重み)でGEMMに寄せる
  3. 固定shapeを作る(x_thr / x_lat)
  4. CUDA Graph capture → replay(ここが支配的)

環境依存で効く(実測推奨)
5) pad_classes=16(あるいは10)
6) F.linear vs mm
7) in-place vs out-of-place ReLU(compileと相性)
8) torch.compile modeの選択


10. “2800万枚/秒”を性能モデルで読む(推測)

hidden=512, pad=16 とすると、1サンプルあたりの主計算量は概算で:

  • 1層目GEMM: 784×512 のMAC → 784*512*2 ≈ 802,816 FLOPs
  • 2層目GEMM: 512×16 のMAC → 512*16*2 = 16,384 FLOPs
  • bias+ReLUは相対的に小さい

合計 ≈ 0.82 MFLOPs / sample

ベースラインの 27.8M samples/s は、

  • 27.8e6 × 0.82e6 ≈ 22.8 TFLOPS 相当の実効

Tesla T4で「FP16 GEMMをTensor Coreで回している」ことと整合します。
逆に言えば、ここまで来ると **“モデルを変えずに更に倍速”**は簡単ではなく、次のボトルネック(カーネル選択・メモリ階層・融合)が見えてきます。


まとめ

  • 2,800万枚/秒級を現実にした本質は、
    「推論をGEMM中心に設計」+「入力も重みもFP16」+「固定shapeでCUDA Graph replay」 の三点セットです(事実として microbench→最終で約2.75x)。
  • v7はさらに 全層CE微調整で精度を 0.9239 → 0.9679 に引き上げ、推論構造は維持したまま“高精度化”を達成しています(事実)。
  • 速度最適化は「理屈より実測」が正しい局面が多く、ノートブックは pad/実装差/compile mode を実測選択する設計になっており、再現性・移植性の面で実務的です。

参考

ここは参考にしたURLを少しだけまとめておきます。

- PyTorch torch.inference_mode
  https://docs.pytorch.org/docs/stable/generated/torch.autograd.grad_mode.inference_mode.html

- PyTorch torch.compile
  https://docs.pytorch.org/docs/stable/generated/torch.compile.html

- PyTorch CUDA Graph API
  https://docs.pytorch.org/docs/stable/generated/torch.cuda.CUDAGraph.html
  https://docs.pytorch.org/docs/stable/generated/torch.cuda.graph.html

- PyTorch Blog: CUDA Graphs
  https://pytorch.org/blog/accelerating-pytorch-with-cuda-graphs/

- NVIDIA: Matrix Multiplication / Mixed Precision / Tensor Core alignment
  (ノートブック内のリンク参照)

Discussion

ピン留めされたアイテム
just for testingjust for testing

とても興味深い内容でした!私もColab/T4で試してみましたので、結果を置いておきます!

--- BENCHMARK ---
Total Time: 0.2304s
Throughput: 86,807,452 samples/sec
Accuracy: 96.91% (Target: ≥96%)
mnist.py
import time
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.datasets as datasets

device = torch.device('cuda')
torch.backends.cudnn.benchmark = True
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cuda.matmul.allow_fp16_reduced_precision_reduction = True

class SmallMLP(nn.Module):
    def __init__(self, hidden_dim=40):
        super().__init__()
        self.net = nn.Sequential(
            nn.Flatten(),
            nn.Linear(784, hidden_dim),
            nn.ReLU(inplace=True),
            nn.Linear(hidden_dim, 10),
        )

    def forward(self, x):
        return self.net(x)

print("Loading data...")

train_ds = datasets.MNIST(root='./data', train=True, download=True)
test_ds  = datasets.MNIST(root='./data', train=False)

X_train = train_ds.data.unsqueeze(1).float().div(255.0).to(device, non_blocking=True)
X_test  = test_ds.data.unsqueeze(1).float().div(255.0).to(device, non_blocking=True)
y_train = train_ds.targets.to(device, non_blocking=True)
y_test  = test_ds.targets.to(device, non_blocking=True)

X_train = X_train.view(-1, 784)
X_test  = X_test.view(-1, 784)

model = SmallMLP().to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
loss_fn = nn.CrossEntropyLoss()

BATCH_SIZE = 2048
EPOCHS = 20

print("Starting training...")
model.train()
t0 = time.time()
indices = torch.arange(X_train.size(0), device=device)

for _ in range(EPOCHS):
    perm = indices[torch.randperm(indices.size(0))]

    for i in range(0, X_train.size(0), BATCH_SIZE):
        idx = perm[i : i + BATCH_SIZE]
        data = X_train[idx]
        target = y_train[idx]

        optimizer.zero_grad(set_to_none=True)
        logits = model(data)
        loss = loss_fn(logits, target)
        loss.backward()
        optimizer.step()

print(f"Training finished in {time.time() - t0:.2f}s")

print("\n--- BENCHMARK ---")

model.half().eval()
X_test_fp16 = X_test.half()

use_graph = True
g = torch.cuda.CUDAGraph()

if use_graph:
    # Warmup and Graph Capture
    static_input = X_test_fp16

    # Warmup model before capture
    for _ in range(5):
        _ = model(static_input)
    torch.cuda.synchronize()

    # Capture the graph
    with torch.cuda.graph(g):
        static_output = model(static_input)

# Warmup the benchmarking loop itself
with torch.no_grad():
    for _ in range(20):
        if use_graph:
            g.replay()
        else:
            _ = model(X_test_fp16)

# Execute Timing
num_runs = 2000
total_samples = X_test_fp16.size(0) * num_runs

torch.cuda.synchronize()
t0 = time.perf_counter()

with torch.no_grad():
    if use_graph:
        for _ in range(num_runs):
            g.replay()
    else:
        for _ in range(num_runs):
            _ = model(X_test_fp16)

torch.cuda.synchronize()
t1 = time.perf_counter()

# --- Results ---
total_time = t1 - t0
throughput = total_samples / total_time
print(f"Total Time: {total_time:.4f}s")
print(f"Throughput: {throughput:,.0f} samples/sec")

with torch.no_grad():
    logits = static_output if use_graph else model(X_test_fp16)
    preds = logits.argmax(dim=1)
    acc = (preds == y_test).float().mean() * 100

print(f"Accuracy: {acc:.2f}% (Target: ≥96%)")
4
just for testingjust for testing

使っているコードは別物になります。わかりづらくて申し訳ありません。先程のコメントの最下部にある"▶︎"をクリックするとコードが展開されるかと思います。

1
just for testingjust for testing

ありがとうございます!ただ、データの転送時間を考慮していないので実用上はもっと遅くなるかと思います…

1