GPU 开销:为什么我们的 MNIST 模型在 CPU 上训练更快
机器学习和 GPU 似乎密不可分——每篇教程、每家云服务商、每份「入门指南」都在把你引向 GPU 实例。直觉很直白:数据进来,几千个核心并行处理,训练飞快。但这里有个前提:那几千个核心真的在被用上——而对小规模的工作负载来说,它们并没有。
我们在一台本地工作站上训练了我们的 MNIST 模型,配置是 Quadro RTX 5000 显卡和 Intel Core i9-10885H 处理器,然后又只用 CPU 跑了一遍。CPU 更快——每轮 4.4 秒对 6.5 秒。本文将解释原因,并展示我们如何在不升级硬件的前提下,靠削减开销把训练时间压到 0.14 秒。
我们先来看训练过程中究竟发生了什么。
拆解一个训练步
每个训练步会走完四个阶段——前向传播、损失、反向传播、权重更新:
# 前向传播
z1 = xb @ w1 + b1 # 矩阵乘法 + 偏置(隐藏层)
a1 = np.maximum(0, z1) # ReLU 激活
z2 = a1 @ w2 + b2 # 矩阵乘法 + 偏置(输出层)
exp_z = np.exp(z2 - z2.max(axis=1, keepdims=True))
probs = exp_z / exp_z.sum(axis=1, keepdims=True) # softmax → 概率
# 损失的梯度
dz2 = probs.copy()
dz2[np.arange(bs), yb] -= 1 # 预测偏离了多少
dz2 /= bs
# 反向传播——计算梯度
dw2 = a1.T @ dz2 # W2 的梯度
da1 = dz2 @ w2.T # 向后流动的梯度
dz1 = da1 * (z1 > 0) # ReLU 的梯度
dw1 = xb.T @ dz1 # W1 的梯度
# 更新权重
w1 -= lr * dw1
b1 -= lr * dz1.sum(axis=0)
w2 -= lr * dw2
b2 -= lr * dz2.sum(axis=0)时间都花到哪儿去了?为了搞清楚,我们用 time.perf_counter() 把每个操作包了起来:
t = time.perf_counter()
z1 = xb @ w1 + b1
timings["fwd: X @ W1 + b1"] += time.perf_counter() - t
t = time.perf_counter()
a1 = np.maximum(0, z1)
timings["fwd: ReLU"] += time.perf_counter() - t
# ...其余每个操作依此类推在 GPU 上运行时,我们必须在各操作之间插入
tf.test.experimental.sync_devices()——否则 GPU 操作会异步排队,计时只能捕捉到派发动作,而不是实际执行。
以下是单个训练步的结果(batch_size=32,对 1000 步取平均,时间单位为微秒 μs——百万分之一秒,越小越好):
| 操作 | CPU (NumPy) | GPU (TF) |
|---|---|---|
| fwd: X @ W1 + b1(图像 × 隐藏层权重 + 偏置) | 162 μs | 541 μs |
| fwd: ReLU(把负值清零) | 9 μs | 173 μs |
| fwd: a1 @ W2 + b2(隐藏层 × 输出层权重 + 偏置) | 19 μs | 478 μs |
| fwd: softmax + 损失(概率 + 误差) | 30 μs | 780 μs |
| bwd: 梯度(每个权重对误差贡献了多少) | 188 μs | 1,271 μs |
| update: W -= lr*dW(调整权重以减小误差) | 345 μs | 1,204 μs |
| 合计 | 751 μs | 4,447 μs |
首先扎眼的一点是:CPU 在每一个操作上都更快——合计 751 μs 对 4,447 μs。本该最受益于 GPU 并行化的操作是矩阵乘法 X @ W1——那是几千个彼此独立的点积,本可以在几千个核心上同时算完,正是 GPU 为之而生的那类工作。
然而在上表中,它在 GPU 上依然更慢:CPU 上 162 μs,GPU 上 541 μs。CPU 在这里快,是因为 NumPy 直接调用 BLAS(Basic Linear Algebra Subprograms)——高度优化的 C/Fortran 例程,会用上 CPU 专有的 SIMD 指令(AVX2 每周期处理 8 个浮点数,FMA 把乘法和加法融合成一条指令)。GPU 更慢,不是因为它的算术更慢,而是因为每个操作都要付出 CUDA 开销(内核启动、显存同步、上下文切换)。
当这些开销比实际计算还费时间,GPU 就算算得更快,最终也会更慢。
我们来看 xb @ w1 + b1 和 a1 @ w2 + b2——两层各自的矩阵乘法——以及 GPU 为何被设计成能把它们做得很快。下面是我们 NumPy 实现中的前向传播:
class HiddenLayer:
def forward(self, x):
self.z = self.W @ x + self.b # 矩阵乘法 + 偏置
self.out = np.maximum(0, self.z) # ReLU 激活
return self.out
class OutputLayer:
def forward(self, x):
self.z = self.W @ x + self.b # 矩阵乘法 + 偏置
exp = np.exp(self.z - np.max(self.z))
self.probs = exp / np.sum(exp) # softmax → 概率
return self.probs每个 @ 都是一次矩阵乘法。对单张图像,self.W @ x 把权重矩阵乘以该图像的像素。128 个神经元中的每一个都有 784 个权重——每个输入像素一个。
每个神经元计算它那 784 个权重与 784 个输入值(图像像素)的点积——即每个神经元 784 次乘法加 783 次加法。结果是 128 个激活值:它们来自该层的 128 个神经元,每个神经元由自己的点积产出一个值。也就是说,一张图像对应 128 个彼此独立的点积:
x (784,) W (128 × 784) result (128,)
┌─────────────┐ ┌──────────────────┐ ┌──────────────┐
│ p1 p2 … p784│ @ │ n1: w1 w2 … w784│ = │ a1 │
└─────────────┘ │ n2: w1 w2 … w784│ │ a2 │
│ ... │ │ ... │
│ n128:w1 w2 … w784│ │ a128 │
└──────────────────┘ └──────────────┘左边是一张 784 像素的图像,中间是 128 个神经元,右边是 128 个激活值。
但 Keras 这类框架并不逐张处理图像——它们把整个批次堆成一个矩阵,一次性乘完。批大小为 32 时,X 是 32 行 × 784 列,同一个权重矩阵在一次操作中产出 32 × 128 = 4096 个激活值:
X (32 × 784) W1 (784 × 128) result (32 × 128)
┌─────────────────┐ ┌──────────────────┐ ┌──────────────────┐
│ img1: p1 … p784│ │ n1 n2 … n128 │ │ img1: a1 … a128 │
│ img2: p1 … p784│ @ │ w w … w │ = │ img2: a1 … a128 │
│ ... │ │ ... ... ... │ │ ... │
│ img32: p1 … p784│ │ w w … w │ │ img32:a1 … a128 │
└─────────────────┘ └──────────────────┘ └──────────────────┘左边是 32 张图像、每张 784 像素,中间是每个神经元 784 个权重,右边是 32 × 128 个激活值。
权重是共享的——这就像你有 32 份相同的模型副本同时跑在 32 张图像上,只不过这是用一次操作完成的,而不是 32 次分开的操作。一次内核启动、一次从显存读取权重矩阵,全部 4096 个结果并行算出。
现在有意思的问题是:这 4096 个点积中,GPU 能同时算多少个? 既然它们彼此完全独立——第 5 张图像在第 20 个神经元上的激活值不依赖任何别的结果——GPU 就可以给每一个分配一个独立线程。
RTX 5000 有 48 个流多处理器(SM),每个含 64 个 CUDA 核心——总共 3072 个核心。每个核心运行在约 1.8 GHz,每周期可做一次乘加,于是理论峰值为:3072 核心 × 每秒 18 亿周期 × 每周期 2 次运算 = 约 11 TFLOPS(每秒万亿次浮点运算)。
我们的矩阵乘法产出 4096 个点积,每个需要 784 次乘加:4096 × 784 = 约 320 万次运算。 GPU 每秒能做 11 万亿次运算,而我们只要 320 万次——那是 0.3 微秒的实际计算,只用到 GPU 能力的 0.00003%。 GPU 用 0.3 μs 算完,然后等大约 540 μs 才轮到下一次内核启动。它只有 0.06% 的时间在干活。
时间实际上花在哪里
在 GPU 上训练神经网络不只有矩阵乘法。每个训练步都是一串操作,而对小模型来说,真正的算术是其中最小的一部分:
1. 内核启动开销。 每个 GPU 操作——矩阵乘法、激活函数、损失计算——都是一个内核(kernel),需要 CPU 去调度并在 GPU 上启动。每次启动都有大约 5–15 微秒的固定开销。对耗时若干毫秒的大矩阵乘法而言,这可以忽略;而对我们这个几微秒就完事的 32×784 × 784×128 小乘法,启动开销可能超过计算本身。仅前向传播就涉及一次矩阵乘法、一次加偏置、一次 ReLU、又一次矩阵乘法、又一次加偏置,再加一次 softmax——在反向传播还没开始之前,至少已有 6 次内核启动。
2. 显存传输延迟。 数据必须从 CPU 内存搬到 GPU 显存(梯度还得搬回来)。这种传输有固定延迟——设置 DMA 传输、跨过 PCIe 总线、发出完成信号所需的时间。批大小 32、输入 784 维时,每批约传输 100 KB。PCIe 总线能跑 32 GB/s,所以纯传输约 3 微秒——但准备阶段的开销是它的 10–20 倍。
3. Python 与框架开销。 Keras/TensorFlow 又加了一层间接。每个操作都要经过 Python、TensorFlow 运行时、XLA 编译(一个在首次运行时优化计算图的编译器)、内存分配和同步。对大操作,这层开销看不见;对小操作,它就是瓶颈。
好消息是,每一类开销究竟占多少时间,我们可以精确测出来。
剖析 CUDA 开销
逐操作的计时显示 GPU 用时更多,而我们算出真正的算术只需 0.3 μs——GPU 大部分时间本该是空闲的。那它把另外 540 μs 究竟花在了什么上?
为了找出答案,我们用了 NVIDIA 的 nsys(Nsight Systems)性能剖析器,它会拦截每一次 CUDA API 调用:
nsys profile -o keras-gpu-profile python mnist-keras.py
nsys stats --force-export=true keras-gpu-profile.nsys-repCUDA 是 NVIDIA 在你的代码与 GPU 硬件之间的软件层。当 TensorFlow 想做矩阵乘法时,它并不直接跟 GPU 对话——它调用「分配显存」「拷贝这些数据」「启动这个内核」之类的 CUDA 函数。每一次这样的调用都要经过驱动,各自都有开销。nsys 剖析器会把它们逐一记录下来,于是我们就能看清时间的去向。
这一次,我们不再计时单个操作,而是剖析了 model.fit() 的整整一轮——全部 1500 步。这一轮在 GPU 上耗时 6.5 秒。时间去向如下:
| 操作 | 开销类型 | 调用次数 | 总耗时 |
|---|---|---|---|
cuCtxSetCurrent(上下文切换) | 内核启动 | 98,168 | 0.87 秒 |
cuEventRecord(计时/同步事件) | 内核启动 | 22,786 | 0.57 秒 |
cuMemcpyDtoHAsync(GPU→CPU 拷贝) | 显存传输 | 6,144 | 0.56 秒 |
cuMemcpyHtoDAsync(CPU→GPU 拷贝) | 显存传输 | 5,678 | 0.27 秒 |
cuGraphLaunch(执行已编译的图) | 内核启动 | 1,875 | 0.16 秒 |
cuLaunchKernel | 内核启动 | 1,080 | 0.02 秒 |
| CUDA 开销合计 | 约 2.7 秒 |
这些 CUDA 条目中没有一项是真正的矩阵乘法——它们全是围绕它的管理性开销。真正的计算(矩阵乘法、ReLU、softmax)发生在 cuGraphLaunch 派发工作之后的 GPU 上,但对我们这些小矩阵而言快到根本排不进显著项。
把所有部分拼起来,这 6.5 秒的一轮是这样分布的:
| 耗时 | |
|---|---|
| CUDA 开销(见上表) | 约 2.7 秒 |
| Python/框架开销(nsys 未捕获) | 约 3.8 秒 |
| 真正的 GPU 算术 | 可忽略 |
| 一轮总计 | 6.5 秒 |
本质上,这 6.5 秒全都是开销。Keras 其实已经优化得相当好——它使用 CUDA Graphs(cuGraphLaunch——1875 次调用,每批一次)把整个前向+反向过程预先编译好,回放时无需逐操作派发。但即便有这项优化,GPU 也始终没机会靠更快的计算把开销赚回来。
作为对比,Keras 在 CPU 上跑同一轮只需 4.4 秒——比 GPU 的 6.5 秒更快,因为完全没有 CUDA 开销。XLA 编译成原生代码,CPU 直接在自己的内存里做算术。
需要强调的是,就算术本身而言并无意外——GPU 确实明显更快,符合预期。我们给核心操作计了时——一次 (32, 784) @ (784, 128) 的矩阵乘法,即一批数据通过第一层:
| 每次矩阵乘法耗时 | |
|---|---|
| GPU | 109 μs |
| CPU | 272 μs |
GPU 快 2.5 倍——但请注意,这 109 μs 里已经包含了启动这一个操作的 CUDA 开销。纯算术大约是 0.3 μs(如前面从 TFLOPS 推算的那样);另外约 108 μs 是这一次内核启动的开销。在 CPU 上,那 272 μs 全是算术——中间没有开销层。
一个完整的训练步包含多次矩阵乘法,外加激活、损失、梯度,以及围绕它们每一项的全部 CUDA 管理。当我们给一个把全部开销都算进去的已编译 train_step 计时,差距就彻底消失了:
| 每步 | |
|---|---|
| GPU | 1.27 毫秒(约 0.1 毫秒计算 + 约 1.2 毫秒开销) |
| CPU | 1.30 毫秒(全是计算,没有开销) |
GPU 算得更快,但其余时间都用在开销上,结果每步的速度大致相当。而在完整的 model.fit() 流程里——包含数据加载、指标和回调——CPU 反而胜出:每轮 4.4 秒对 6.5 秒。
一个有趣的细节:如果你在训练时运行 nvidia-smi,可能会看到 GPU 利用率是 90–100%。看起来 GPU 忙得不可开交——那它为什么还比 CPU 慢?
因为 nvidia-smi 报告的是至少有一个内核在 GPU 上运行的时间占比——而不是有多少核心在工作。如果一个个小内核首尾相接地启动、中间没有空隙,它就会显示约 100% 的利用率,尽管在任一时刻绝大多数核心都在闲着。正如前面所见,我们的模型只用到 GPU 能力的 0.00003%。
让单次训练跑得更快
既然理解了开销的来源,该怎么削减它?有好几个切入角度:减少步数(加大批大小)、消除 CUDA 开销(在 CPU 上跑)、消除框架开销(不用 TensorFlow)、把整个步骤编译成一次操作(JAX 的 JIT),或者天真地搬到 GPU 上(CuPy——剧透:更糟)。 这五条我们都试了。
1. 加大批大小
如果开销是按步付的,最显然的办法就是:少走几步。更大的批意味着每步处理更多样本,因此同一轮所需的步数更少——而每步只需缴一次开销税,与批大小无关。训练样本 48000 个(6 万减去 20% 的验证集)、批大小 32 时,一轮是 48000 / 32 = 1500 步,每步是在一个批次上完整走一遍前向 + 反向 + 更新。批大小取 4096 时,只有 11 步。
更大的批也让 GPU 每步有更多活干——如前所述,批大小 32 只让 48 个 SM 中的约 3 个保持活跃,而批大小 2048 能让约 42 个忙起来。所以加大批大小有两重好处:每轮需要缴开销的步数更少,且每步真正在工作的 GPU 核心更多。
不过,批大小不能不加分辨地往上加——它直接影响模型准确率。更大的批带来更平滑但更不频繁的梯度更新,可能导致泛化变差。合适的批大小需要你针对自己的模型试出来。
2. Keras 跑在 CPU 上——跳过 CUDA 开销
如前所见,仅仅禁用 GPU(tf.config.set_visible_devices([], 'GPU'))就消掉了 2.7 秒的 CUDA 开销。XLA 编译成原生代码,CPU 直接做算术。
3. 纯 NumPy——把框架开销也一并跳过
再往前一步,我们可以完全不用 TensorFlow。纯 NumPy 的训练循环直接调用 BLAS 例程,操作之间没有框架派发开销。
4. JAX 的 JIT——把整个步骤编译起来
JAX 的 jit 编译器会追踪整个训练步,并把它编译成一个经过优化的原生函数。不再由 Python 逐个派发操作,编译后的函数在一次融合调用中把它们全部执行完,开销几乎为零。
5. CuPy——要是干脆把 NumPy 搬到 GPU 上呢?
我们也试了这个天真的路子:把 import numpy as np 换成 import cupy as cp,同样的代码在 GPU 上跑。由于我们的实现是在 Python 循环里逐个样本处理,每个微小操作(一次 128 元素的向量加法、一次 10 元素的 softmax)都变成一次独立的 GPU 内核启动。结果是:5 轮耗时 443 秒——比 NumPy 在 CPU 上慢了 6 倍以上。不重新思考访问模式就把代码搬上 GPU,只会更糟,不会更好。
把这些合起来看
我们在不同批大小下测试了方案 1–4(CuPy 太慢,未纳入)。每个单元格是一轮的耗时:
| 批大小 | Keras GPU | Keras CPU | 纯 NumPy | JAX (CPU, JIT) |
|---|---|---|---|---|
| 32 | 6.26 秒 | 11.31 秒 | 2.71 秒 | 2.37 秒 |
| 128 | 2.74 秒 | 2.21 秒 | 4.70 秒 | 1.03 秒 |
| 512 | 2.64 秒 | 1.80 秒 | 1.30 秒 | 0.64 秒 |
| 1024 | 2.47 秒 | 1.43 秒 | 0.87 秒 | 0.63 秒 |
| 2048 | 2.32 秒 | 1.42 秒 | 0.84 秒 | 0.48 秒 |
| 4096 | 2.32 秒 | 1.39 秒 | 1.03 秒 | 0.14 秒 |
横着读每一行,看到的是换方案(每步开销更少)的效果;竖着读每一列,看到的是加大批(每轮步数更少)的效果。两种效果会叠加。
Keras 在 GPU 上触底于约 2.3 秒,无论批大小如何——CUDA 开销存在一个固定的地板,加大批也消不掉。
Keras 在 CPU 上从 11.3 秒降到 1.4 秒——没有 CUDA 开销,但 TensorFlow 自身的框架开销划出了另一条地板。
纯 NumPy 在批大小 2048 时达到 0.84 秒——没有框架,只有 BLAS 调用。到 4096 时因内存压力又慢了回去。
JAX 在批大小 4096 时达到 0.14 秒——比 Keras 在 GPU 上快 45 倍。JIT 编译把整个步骤融合成一次原生调用,开销几乎为零。
教训是:通往更快训练的路不是更大的硬件,而是更少的开销。 表中每种方案都剥掉了一层开销,而更大的批则减少了你为剩余开销付费的次数。
用并行化做超参数搜索
上面这些办法让单次训练更快。但当你做超参数搜索时——试 5 个学习率,或 4 种架构——每次运行都是完全独立的。它们不共享权重、梯度或状态。
JAX 的 vmap(向量化 map)能在一次前向和反向传播中同时训练全部 5 个模型。
这与 Keras 的批处理不同,后者是把许多图像送过一个模型。
vmap 是把同一批图像送过许多模型——每个模型有自己的权重和学习率——在一次融合操作中完成。在底层,XLA 会把它编译成单个内核:一次矩阵乘法处理全部 5 次前向传播,另一次处理全部 5 次反向传播。没有 Python 循环,也没有按模型计的开销。
代价是:所有模型必须具有相同的架构(相同的参数形状),而且训练循环得用 JAX 重写。完整代码如下:
import jax
import jax.numpy as jnp
from jax import vmap, jit, grad, random
def init_params(key):
# 与之前相同的模型:784→128→10,随机权重
k1, k2 = random.split(key)
w1 = random.normal(k1, (784, 128)) * jnp.sqrt(2.0 / 784)
b1 = jnp.zeros(128)
w2 = random.normal(k2, (128, 10)) * jnp.sqrt(2.0 / 128)
b2 = jnp.zeros(10)
return (w1, b1, w2, b2)
def loss_fn(params, x, y):
# 前向传播 + 交叉熵损失——与我们 NumPy 版本相同的数学
w1, b1, w2, b2 = params
h = jnp.maximum(0, x @ w1 + b1) # 隐藏层 + ReLU
logits = h @ w2 + b2 # 输出层
log_probs = logits - jnp.log(jnp.sum(jnp.exp(logits), axis=-1, keepdims=True))
return -jnp.mean(log_probs[jnp.arange(y.shape[0]), y])
def sgd_step(params, x, y, lr):
# 一个训练步:计算梯度,更新权重
grads = grad(loss_fn)(params, x, y) # JAX 自动对 loss_fn 求导
return tuple(p - lr * g for p, g in zip(params, grads))
# 5 个学习率、5 套权重,同时训练
lr_array = jnp.array([0.001, 0.01, 0.1, 1.0, 10.0])
batched_step = jit(vmap(sgd_step, in_axes=(0, None, None, 0)))
batched_params = batched_step(batched_params, x_batch, y_batch, lr_array)对更简单的情形,Python 的 multiprocessing.Pool 也能达到同样目的——5 个独立进程,各训练一种配置:
import multiprocessing as mp
def train_one_config(args):
# 用一个学习率训练一个模型——在自己的进程中运行
lr, X_train, y_train = args
model = keras.Sequential([
keras.layers.Dense(128, activation="relu", input_shape=(784,)),
keras.layers.Dense(10, activation="softmax"),
])
model.compile(optimizer=keras.optimizers.SGD(learning_rate=lr),
loss="sparse_categorical_crossentropy")
history = model.fit(X_train, y_train, epochs=10, batch_size=32,
validation_split=0.2, verbose=0)
return lr, history.history
with mp.Pool(5) as pool:
results = pool.map(train_one_config,
[(lr, X_train, y_train) for lr in [0.001, 0.01, 0.1, 1.0, 10.0]])几点要带走的结论
- 对小模型,CPU 比 GPU 更快。 我们的 MNIST 模型(784→128→10,约 10.1 万参数)在 CPU 上每轮 4.4 秒,在 GPU 上 6.5 秒。GPU 的算术更快(每次矩阵乘法 109 μs 对 272 μs),但 CUDA 开销——内核启动、显存拷贝、上下文切换——每轮要多出 2.7 秒,而 CPU 根本没有这笔开销。
- GPU 几乎没被用上。 我们的矩阵乘法只用到 RTX 5000 能力的 0.00003%。在 batch_size=32 时,48 个 SM 中只有约 3 个活跃。GPU 用 0.3 μs 算完,然后等约 540 μs 才轮到下一次内核启动。
nvidia-smi的利用率会误导人。 它报告的是「有任何内核在运行」的时间占比,而不是核心占用率。我们的模型显示约 100% 利用率,却只用了 GPU 算力的 0.00003%。- 削减开销胜过更大的硬件。 从 Keras + GPU 换到 JAX JIT + CPU 并把 batch_size 设为 4096,我们把每轮从 6.5 秒压到 0.14 秒——同一台机器上 45 倍加速,完全不需要 GPU。
- 在工作负载真正需要之前,别急着上 GPU。 对参数少于约 50 万、批大小低于 256 的模型,一颗快的 CPU 会同时更便宜也更快。
同样的规律在生产环境中还有个著名的例子——word2vec:每个训练步只触及两个嵌入矩阵中的少数几行,做的是 10⁴ 量级的乘加——每个样本的计算量比上面 MNIST 的前向传播少了大约一千倍。word2vec 是一种以内存访问和随机行查找为主导的负载,而不是稠密矩阵乘法负载,而这恰恰就是在 GPU 上吃亏、在 CPU 上占优的那种形态。事实上的标准 word2vec 库 Gensim 把这一点用到了极致:Cython 写的内层循环、Hogwild! 式的无锁 CPU 线程、全程不用 GPU——在从语料训练一张 (V, d) 查找表这件事上,它的挂钟时间至今仍胜过 GPU 实现,而这已是 CUDA 开始主导深度学习十多年之后的事了。word2vec 那篇文章会详细讲清其中的原因。