ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

Tensor Core原理与实践:大模型矩阵乘法加速的底层逻辑

Tensor Core原理与实践:大模型矩阵乘法加速的底层逻辑 大模型训练和推理为什么总在抢 GPU抛开显存容量不谈最直接的原因是矩阵乘法的计算量太大了。Transformer 里从 embedding 到注意力再到 FFN几乎每一步都在做矩阵乘法一个参数规模很大的模型在训练时算力主要消耗在 GEMMGeneral 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 方向每次前进BKfor (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..i0BM, j0..j0BN 区域 } }分块之后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 核心最常用的数学操作是 FMAfused 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 c16x16或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 CoreTF32 开关、基准和 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, devicecuda, dtypetorch.float32) B torch.randn(K, N, devicecuda, dtypetorch.float32) def bench(fn, name, repeat20): # 预热触发第一次分配和 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_bycuda_time_total, row_limit10))再运行一次开启 TF32 的情况对比 CUDA kernel 耗时变化。不过有一点需要强调cuBLAS 的 kernel 名称在不同 CUDA 版本中并不稳定。仅凭名字里带有sgemm、gemm并不能百分百判断 Tensor Core 是否参与。更可靠的方式是使用 Nsight Computencu --set full python profile_script.py如果环境允许 Profiling 计数器ncu会显示计算管道利用率、内存吞吐、Tensor Core 相关指标。它比“看时间快了多少”更能说明问题。4. 如果想写 CUDA Kernel矩阵分块要如何踩准 Tensor CorePyTorch 层已经封装了很多细节。但理解 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::fragmentwmma::matrix_a, 16, 16, 16, __half, wmma::row_major a_frag; wmma::fragmentwmma::matrix_b, 16, 16, 16, __half, wmma::row_major b_frag; wmma::fragmentwmma::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这条指令的含义通常包括m16n8k8A 是 16x8B 是 8x8C/D 是 16x8。row.colA 按 row-major 排列B 按 col-major 排列。f32表示累加寄存器是 FP32。tf32.tf32表示输入
返回列表