大模型训练和推理为什么总在抢 GPU?抛开显存容量不谈,最直接的原因是矩阵乘法的计算量太大了。Transformer 里从 embedding 到注意力再到 FFN,几乎每一步都在做矩阵乘法;一个参数规模很大的模型在训练时,算力主要消耗在 GEMM(General Matrix Multiply)上。NVIDIA 为这类计算准备的专用硬件就是 Tensor Core。Tensor Core 不是像普通 CUDA 核心那样一次只处理若干标量乘加,而是让一个线程束协作完成一个小块矩阵的乘加。下面从矩阵分块入手,把 Tensor Core 的加速原理、PyTorch 中的验证方法、手写 CUDA 时的调用思路,以及常见的“加速失效”场景讲清楚。
1. 大模型的计算压力先从矩阵乘法讲起
矩阵乘法不是大模型独有的问题,但大模型把矩阵乘法的规模推到了新的高度。只有先理解矩阵乘法在 Transformer 中出现在哪些地方、为什么计算量巨大,才能理解 Tensor Core 为什么重要。
1.1 Transformer 的矩阵乘法分布
在典型 Transformer 结构中,训练和推理阶段的核心算子可以简化为一张表:
| 计算环节 | 简化的矩阵表达 | 说明 |
|---|---|---|
| QKV 投影 | X * W_qkv | 输入序列向量和权重矩阵相乘,得到 Q、K、V |
| 注意力分数 | Q * K^T | 每个位置与其它位置的相似度 |
| 注意力输出 | attn_weights * V | 对 Value 按注意力权重加权 |
| FFN 第一层 | Y * W1 | 升维线性变换 |
| FFN 第二层 | Z * W2 | 降维回原始维度 |
| Token 输出层 | hidden * W_cls | 将 hidden state 映射到词表或标签空间 |
只看“计算次数”,线性层和注意力部分的矩阵乘会占到绝大多数算力。即使引入 Flash Attention、KV Cache 等优化,模型仍然不可能绕开线性层的矩阵乘法。
1.2 矩阵乘法为什么不能靠最简单三重循环
最直观的矩阵乘法写法如下:
for (int i = 0; i < M; ++i) { for (int j = 0; j < N; ++j) { float sum = 0.0f; for (int k = 0; k < K; ++k) { sum += A[i * K + k] * B[k * N + j]; } C[i * N + j] = sum; } }这段代码在三重循环层面没有任何错误,但性能很差。原因有两类:
- B 矩阵的访问不连续。在最内层,每次都会跳到 B 的第
k行,缓存利用很差。 - 标量循环很难并行。GPU 擅长让大量线程同时做相同工作,而不是让一个线程串行执行成千上万次迭代。
正确的做法不是把循环“拍平”,而是把矩阵切成块,让每个线程或每个线程束负责一块连续数据。
1.3 分块矩阵乘法如何增加计算密度
矩阵分块的基本思想是:与其让一个线程负责一个输出元素,不如让一组线程负责一个输出小块。
假设输出 C 被分成BM行、BN列的小块,K 方向每次前进BK:
for (int i0 = 0; i0 < M; i0 += BM) { for (int j0 = 0; j0 < N; j0 += BN) { float acc[BM][BN] = {}; for (int k0 = 0; k0 < K; k0 += BK) { // 将 A 的一个 BM*BK 分块拷贝到 shared memory // 将 B 的一个 BK*BN 分块拷贝到 shared memory // 对 A 分块和 B 分块做小矩阵乘法,累加到 acc } // 把 acc 写回 C 的 i0..i0+BM, j0..j0+BN 区域 } }分块之后,A、B 的数据在片上可以被重复使用。一次BM * BK与BK * BN的小块乘法,计算量是2 * BM * BN * BK,需要从主存读取的数据量约为BM * BK + BK * BN。当BM和BN都变大时,一次分块计算中能摊薄的数据搬运成本就越多,计算密度也就越高。
这里有一个关键认知:GPU 高层应用看的是“谁在计算”,底层看的是“内存搬运和寄存器复用”。Tensor Core 的分块设计就是为了让这一层复用关系变得可控。
2. Tensor Core 的加速原理:一次完成一个矩阵小块的乘加
Tensor Core 并不是把所有矩阵乘法都变成“神奇硬件”。准确地说,它把矩阵乘法的基本操作从“单个标量乘加”升级成了“一个线程束协同完成矩阵块乘加”。
2.1 普通 FMA 与 Tensor Core 的差别
普通 CUDA 核心最常用的数学操作是 FMA(fused multiply-add),即d = a * b + c。这个操作一次只处理一组标量。
Tensor Core 处理的是矩阵乘加:
D = A * B + C这不是一个线程能独立完成的操作,而是一个线程束级别操作。通常由 32 个线程协作提供 A 块、B 块和 C 块的数据,再由 SM 中的 Tensor Core 单元执行一次矩阵乘加。
可以用表格简单对比:
| 对比项 | 普通标量核心 | Tensor Core |
|---|---|---|
| 最小计算单元 | 一个线程做一个标量 FMA | 一个 warp 协作做一个矩阵块乘加 |
| 典型数据形状 | a * b + c | 16x16或16x8等矩阵块 |
| 计算类型 | FP32、FP64 等 | FP16、BF16、TF32、INT8 等 |
| 设计目标 | 通用线程计算 | 高密度矩阵乘加吞吐 |
| 程序使用方式 | 普通 CUDA 指令 | WMMA、mma.sync、cuBLAS/CUTLASS 等 |
从 Volta 架构开始,NVIDIA GPU 引入了 Tensor Core;后续 Turing、Ampere、Hopper 等架构不断更新数据类型和矩阵块大小。不同架构支持的 tile 形状不一定相同,常见 PTX 资料里能看到m16n8k4、m16n8k8、m16n8k16、m16n16k16等组合。
2.2 一个 warp 如何完成一次矩阵块乘法
可以先看一个示意:
A tile: 16 行 x 8 列 B tile: 8 行 x N 列 C tile: 16 行 x N 列 C16xN += A16x8 * B8xN实际硬件并不会把矩阵的每个元素都放到同一个线程里。32 个线程各自持有 fragment 的一部分,数据分布规则由硬件定义。写 CUDA 程序时,普通开发者不一定需要直接理解每个线程寄存器里放了哪些元素,但必须知道以下几个事实:
- 数据不能随意地从任意内存地址丢给 Tensor Core,需要符合连续布局和对齐要求。
- 数据要先放进寄存器或 shared memory,再由 warp 级指令去消费。
- 参与同一个 fragment 的线程必须都在同一个 warp 内,并且执行同一段
mma_sync或load_matrix_sync代码。
一个容易误解的地方是:很多人以为 Tensor Core 一次会直接吞掉整张大矩阵。实际上,GEMM 库会把大矩阵继续切分成很多小 tile,逐块送入 Tensor Core。外层是传统的分块调度,内层才是 Tensor Core 的mma指令。
2.3 分块让数据复用而不是让核心空转
GPU 的算力峰值很高,但数据从显存搬到片上需要时间。Tensor Core 能跑多快,不只看它每秒能做多少次矩阵乘加,还要看数据是否来得及从上一级存储搬到寄存器。
分块的作用可以从“重用一个输入元素”的角度看。一个 16x16 的输出 tile,需要读取 A 的一个 16xK 片段和 B 的一个 Kx16 片段。如果 K 方向的累加很长,A 片段和 B 片段会在寄存器中反复参与计算。这样每个加载进来的元素都做了多次计算,而不是加载一次只做一次乘加。
如果没有这种分块复用,就会陷入“内存受限”。即使 Tensor Core 理论算力再高,SM 也只能等着数据从显存或 L2 返回,最终测出来的时间并没有显著下降。
3. PyTorch 里看 Tensor Core:TF32 开关、基准和 Profiler
很多大模型开发者不会直接写 CUDA Kernel,而是在 PyTorch 中调用Linear、matmul或注意力算子。此时 Tensor Core 是否参与计算,取决于框架版本、cuBLAS 策略、输入精度和形状。
3.1 确认显卡、CUDA 与 TF32 开关状态
先写一段环境检查脚本:
import torch print("cuda available:", torch.cuda.is_available()) print("gpu:", torch.cuda.get_device_name(0)) print("torch cuda:", torch.version.cuda) print("matmul allow_tf32:", torch.backends.cuda.matmul.allow_tf32) print("cudnn allow_tf32:", torch.backends.cudnn.allow_tf32) print("device capability:", torch.cuda.get_device_capability(0))输出大致类似:
cuda available: True gpu: NVIDIA GeForce RTX 4090 torch cuda: 12.1 matmul allow_tf32: False cudnn allow_tf32: True device capability: (8, 9)这里有两个开关要区分清楚:
torch.backends.cuda.matmul.allow_tf32:控制矩阵乘法是否允许把 FP32 输入转成 TF32。torch.backends.cudnn.allow_tf32:控制 cuDNN 卷积相关算子是否允许 TF32。
大模型推理时,有些场景会直接使用 FP16 或 BF16;训练时很多人也会用混合精度。但如果你在 FP32 精度下做矩阵乘,希望借助 Ampere 及以上架构的 Tensor Core,就必须确认matmul.allow_tf32的状态。
3.2 用同一块 GPU 对比 FP32、TF32 和 FP16
下面做一个最小验证。使用 4096 维度的矩阵乘法,先做预热,再计时:
import torch import time torch.manual_seed(0) M = K = N = 4096 A = torch.randn(M, K, device="cuda", dtype=torch.float32) B = torch.randn(K, N, device="cuda", dtype=torch.float32) def bench(fn, name, repeat=20): # 预热,触发第一次分配和 kernel 编译 for _ in range(5): fn() torch.cuda.synchronize() start = time.perf_counter() for _ in range(repeat): fn() torch.cuda.synchronize() avg_ms = (time.perf_counter() - start) / repeat * 1000 print(f"{name:12s}: {avg_ms:.2f} ms") # 保持默认,通常不允许 TF32 bench(lambda: A @ B, "fp32 default") # 打开 TF32 torch.backends.cuda.matmul.allow_tf32 = True bench(lambda: A @ B, "fp32 tf32") # 使用 FP16 A16 = A.half() B16 = B.half() bench(lambda: A16 @ B16, "fp16")运行后会发现 FP16 明显更快;TF32 相比关闭时往往也会快不少,但不同显卡、不同矩阵形状下结论会变化。
再比较数值差异:
torch.backends.cuda.matmul.allow_tf32 = False C_fp32 = A @ B torch.backends.cuda.matmul.allow_tf32 = True C_tf32 = A @ B C_fp16 = (A16 @ B16).float() print("fp32 vs tf32 max diff:", (C_fp32 - C_tf32).abs().max().item()) print("fp32 vs fp16 max diff:", (C_fp32 - C_fp16).abs().max().item())实际输出会因矩阵内容和硬件不同而不同,但通常能看到 TF32 和 FP16 与 FP32 之间存在一定误差。误差不是“bug”,而是精度截断的预期结果。
注意:不要只验证程序能跑通。矩阵乘法的验证必须包含“关了开关”和“开了开关”的差异,否则你可能并不知道模型里到底用的是什么精度。
3.3 用 Profiler 看底层 Kernel 和硬件占用
如果只比较时间,还不够严谨。可以通过 PyTorch Profiler 查看 CUDA Kernel 名称:
from torch.profiler import profile, ProfilerActivity torch.backends.cuda.matmul.allow_tf32 = False with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: for _ in range(5): C = A @ B print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))再运行一次开启 TF32 的情况,对比 CUDA kernel 耗时变化。
不过有一点需要强调:cuBLAS 的 kernel 名称在不同 CUDA 版本中并不稳定。仅凭名字里带有sgemm、gemm,并不能百分百判断 Tensor Core 是否参与。更可靠的方式是使用 Nsight Compute:
ncu --set full python profile_script.py如果环境允许 Profiling 计数器,ncu会显示计算管道利用率、内存吞吐、Tensor Core 相关指标。它比“看时间快了多少”更能说明问题。
4. 如果想写 CUDA Kernel,矩阵分块要如何踩准 Tensor Core
PyTorch 层已经封装了很多细节。但理解 Tensor Core 的更好方式是写一个分块 CUDA Kernel,即使只是跑通一个 16x16 的输出 tile,也能建立非常具体的直觉。
4.1 不要从 global memory 直接做 mma
很多初学者会试图写类似这样的逻辑:把大矩阵的某个元素直接传给mma_sync。但 Tensor Core 指令的数据来自寄存器或 shared memory,不可能从全局内存逐元素读取。
一个正常的分块 Kernel 流程是:
- 当前 block 从 global memory 读取 A、B 的一个分块到 shared memory。
__syncthreads()同步。- 从 shared memory 用
load_matrix_sync加载到 WMMA fragment。 - 对 K 方向累加多次后,用
store_matrix_sync把结果写回 shared memory。 - 再把结果从 shared memory 写回 global memory。
共享内存在这里的作用是“中转站”,它把不连续或不齐的内存访问整理成 Tensor Core 能消费的连续布局。
4.2 使用 WMMA API 完成最小矩阵块累加
CUDA 提供nvcuda::wmmaAPI,可以隐藏一部分底层寄存器布局。核心片段如下:
#include <mma.h> #include <cuda_fp16.h> using namespace nvcuda; // 假设当前 CUDA Kernel 已经在一个 warp 中执行 wmma::fragment<wmma::matrix_a, 16, 16, 16, __half, wmma::row_major> a_frag; wmma::fragment<wmma::matrix_b, 16, 16, 16, __half, wmma::row_major> b_frag; wmma::fragment<wmma::accumulator, 16, 16, 16, float> c_frag; // A_tile、B_tile 指向 shared memory 中已加载好的矩阵块 wmma::fill_fragment(c_frag, 0.0f); wmma::load_matrix_sync(a_frag, A_tile, lda); wmma::load_matrix_sync(b_frag, B_tile, ldb); wmma::mma_sync(c_frag, a_frag, b_frag, c_frag); wmma::store_matrix_sync(C_tile, c_frag, ldc, wmma::mem_row_major);这段代码不是完整 Kernel,但已经能看出 Tensor Core 的使用模式:先定义 fragment,再加载矩阵块,然后执行一次矩阵乘累加,最后存回 C 块。
实际生产 Kernel 中会继续展开成 K 方向多层循环,控制每个线程对应多个输出 tile,并用双缓冲隐藏 shared memory 加载延迟。这才是高性能 GEMM Kernel 的核心复杂度所在。
4.3 PTX 层还有更底层的 mma.sync
如果想更深一层,可以看 PTX 指令。例如 TF32 相关的矩阵乘指令可能长这样:
mma.sync.aligned.m16n8k8.row.col.f32.tf32.tf32.f32这条指令的含义通常包括:
m16n8k8:A 是 16x8,B 是 8x8,C/D 是 16x8。row.col:A 按 row-major 排列,B 按 col-major 排列。f32表示累加寄存器是 FP32。tf32.tf32表示输入