1. 从标题拆解Rubin FP4 GEMM 到底在讲什么第一次看到“Rubin FP4 GEMM Overview”这个标题很多人会愣一下Rubin 是什么FP4 又是什么GEMM 我倒是听过但跟 Rubin 放一起就有点陌生了。我先把这三个词拆开讲清楚后面所有内容都围绕它们展开。Rubin 是 NVIDIA 下一代数据中心 GPU 架构的代号接在 Blackwell 后面。Blackwell 这一代大家已经比较熟了RTX 50 系消费卡、B200、GB200 这些产品都基于它。Rubin 则是再往后的规划面向的是更大规模的 AI 训练和推理场景。FP4 是一种数值精度格式全称是 4-bit Floating Point也就是用 4 个比特来表示一个浮点数。GEMM 是 General Matrix Multiply 的缩写通用矩阵乘法深度学习里绝大部分计算量都花在这上面——卷积可以展开成 GEMM注意力机制里的 QK^T 和 PV 也是 GEMM全连接层更是标准的 GEMM。把这三个词连起来Rubin FP4 GEMM 说的就是在 Rubin 这代硬件上如何用 FP4 精度高效地完成矩阵乘法运算。这件事为什么值得单独拿出来讲因为 FP4 相比 FP16、BF16 甚至 FP8位宽又砍了一半意味着同样的显存能装下更多参数、同样的带宽能喂给计算单元更多数据。但代价也很明显4 个比特能表示的数值范围非常有限量化误差会显著增大。所以 Rubin FP4 GEMM 的核心挑战不是“能不能算”而是“怎么算得又快又准”。这篇文章适合谁看如果你在做大模型推理优化、在写 CUDA kernel、在用 CUTLASS 做自定义算子或者你只是对 NVIDIA 新一代硬件的数值格式演进感兴趣那这篇内容应该能给你一些可参考的东西。我会尽量把原理讲透同时给出可以上手操作的思路和参数选择的依据而不是停留在概念层面。需要提前说明的是Rubin 架构目前公开的细节还比较有限很多内容是基于 Blackwell 已有的 FP4 实现和 CUTLASS 的公开接口做的合理推演。我会在涉及推测的地方明确标注避免把猜测当成事实来写。2. 为什么 FP4 值得单独做一代 GEMM 优化2.1 从 FP16 到 FP8 再到 FP4 的精度演进逻辑要理解 FP4 GEMM 的意义得先看清楚精度演进的脉络。早期训练和推理基本都用 FP32后来发现 FP16 和 BF16 在大多数场景下精度足够而且计算吞吐翻倍、显存占用减半于是混合精度训练成了标配。到了 Hopper 这一代FP8 开始被引入E4M3 和 E5M2 两种格式分别面向不同的数值范围需求。Blackwell 进一步把 FP4 推到了台前尤其是配合微缩放Microscaling简称 MX格式使用。这个演进背后的驱动力很直接大模型的参数量在指数级增长而显存带宽和容量的增长是线性的。你不可能无限堆 HBM所以必须让每个参数占用的比特数降下来。FP16 每个参数 16 比特FP8 降到 8 比特FP4 再降到 4 比特。理论上同样的 80GB 显存用 FP16 能放 400 亿参数用 FP4 就能放 1600 亿参数。这个差距在推理场景下是决定性的——能不能单卡放下一个模型往往就取决于此。但精度下降不是没有代价的。FP4 只有 4 个比特其中还要分给符号位和指数位尾数位所剩无几。以常见的 E2M1 格式为例1 位符号、2 位指数、1 位尾数。它能表示的数值只有有限的几个离散值比如 0、0.5、1、1.5、2、3、4、6 这些具体取决于指数偏移。这意味着量化后的权重和激活值会被强行“吸附”到这些离散点上误差不可避免。2.2 微缩放格式如何缓解 FP4 的精度问题单纯用 FP4 做量化精度损失会大到无法接受。所以实际方案里通常采用分块量化的思路把一个大矩阵切成若干小块每一块共享一个缩放因子scale factor。这个缩放因子用较高精度比如 FP8 或 FP16存储块内的元素则用 FP4 表示。计算时先把 FP4 元素乘上对应的缩放因子还原到较高精度再做累加。这就是 MX 格式的核心思想。MXFP4 的块大小通常是 32 个元素共享一个 E8M0 的缩放因子8 位指数、0 位尾数只能表示 2 的幂次。这样做的好处是块内的数值动态范围被缩放因子“对齐”了FP4 只需要在这个相对窄的范围内做精细表示。相比全局量化分块量化的精度损失小得多。在 GEMM 里这个缩放操作需要融入计算流程。假设 A 矩阵是 M×KB 矩阵是 K×NA 和 B 都按 K 维度分块因为累加是沿着 K 方向进行的。每个块有自己的缩放因子计算时需要把 A 的块缩放因子和 B 的块缩放因子相乘再应用到部分和的累加上。这个额外的乘法在硬件里通常有专门的支持不会拖慢主计算流水线。2.3 Rubin 相比 Blackwell 在 FP4 GEMM 上可能的变化Blackwell 已经支持 FP4 GEMM那 Rubin 会带来什么变化从公开信息推断可能有几个方向。第一是张量核心的 FP4 吞吐进一步提升比如从 Blackwell 的某个数值再翻倍。第二是缩放因子的处理更加硬件化减少软件层面的开销。第三是可能引入更灵活的块大小或混合精度模式让开发者能针对不同层选择不同的量化策略。还有一个值得关注的点是 CUTLASS 的支持节奏。CUTLASS 是 NVIDIA 官方的 CUDA 模板库专门用来写高性能 GEMM 和卷积。每一代新架构的 FP4 支持通常都是先在 CUTLASS 里落地然后才逐步进入 cuBLAS、TensorRT 这些上层库。所以如果你想第一时间用上 Rubin 的 FP4 GEMM盯紧 CUTLASS 的 release note 和 example 是最靠谱的路径。3. FP4 GEMM 的核心技术细节与实操要点3.1 数据布局为什么 K 维度分块是主流选择在 FP4 GEMM 里数据布局直接决定了 kernel 的性能和精度。前面提到缩放因子是按块共享的块怎么切、切多大都有讲究。目前主流方案是沿着 K 维度分块块大小为 32。为什么是 K 维度而不是 M 或 N因为 GEMM 的累加是沿着 K 方向进行的。C[i][j] sum over k of A[i][k] * B[k][j]。如果你把 K 切成块每个块内的部分和可以独立计算最后再累加。缩放因子跟着 K 块走意味着每个部分和在累加前先乘上对应的缩放因子。这样做的数学等价性最干净不会引入额外的交叉项。如果按 M 或 N 分块缩放因子会作用在输出矩阵的不同行或列上累加时就需要对每个输出元素单独处理缩放反而更麻烦。所以 K 维度分块是自然选择也是 CUTLASS 里 MXFP4 GEMM 的默认布局。实际操作中A 矩阵通常是行主序row-majorB 矩阵是列主序column-major这样在加载时能保证连续内存访问。FP4 数据在内存里是打包存储的每字节存两个 FP4 值。加载到寄存器后需要解包这个解包操作在 Blackwell 和 Rubin 上都有对应的硬件指令支持比如cvt或者专门的ldmatrix变体。3.2 缩放因子的计算与存储开销缩放因子虽然只占很少的存储空间但它的计算和加载不能忽视。以块大小 32 为例一个 K4096 的矩阵每个 M×K 的行需要 4096/32 128 个缩放因子。如果缩放因子用 FP8 存储那就是 128 字节相比 FP4 数据本身的 4096/2 2048 字节占比约 6%。这个开销在显存带宽紧张的场景下是需要考虑的。更关键的是缩放因子的加载模式。在 GEMM kernel 里A 和 B 的缩放因子需要和对应的数据块一起加载到共享内存或寄存器。如果缩放因子的访问模式不连续就会产生额外的内存事务。CUTLASS 的做法通常是把缩放因子单独放在一块内存区域用专门的加载指令比如ldmatrix的缩放因子版本来读取保证合并访问。还有一个细节是缩放因子的精度。E8M0 只能表示 2 的幂次这意味着缩放因子本身也有量化误差。如果你的数据分布比较均匀这个误差可以接受但如果某些块的动态范围特别大E8M0 可能不够用。这时候可以考虑用 FP16 或 BF16 来存缩放因子代价是存储开销翻倍。具体选哪种需要根据你的模型权重分布来实测。3.3 累加精度FP32 累加器为什么不能省FP4 的乘法结果精度很低但累加必须用 FP32。这是 FP4 GEMM 能work的前提。原因很简单假设你用 FP4 做累加K4096 的情况下累加 4096 个低精度数误差会累积到无法接受的程度。而 FP32 累加器有 24 位尾数能保证在累加过程中不丢失有效精度。在硬件层面张量核心的 FP4 模式通常内置 FP32 累加器。你写 kernel 的时候不需要手动管理累加精度但需要确保输出矩阵是 FP32 或 FP16 格式。如果最终输出要转成 FP16那在写回之前做一次转换即可。这里有个容易踩的坑有些开发者为了省显存想把累加器也设成 FP16。实测下来在 K 较大的情况下比如超过 1024FP16 累加的误差会明显影响最终结果尤其是在 attention 的 softmax 之前那一步。所以除非你的 K 很小否则老老实实用 FP32 累加。3.4 CUTLASS 中 FP4 GEMM 的接口调用要点CUTLASS 的 FP4 GEMM 通常通过cutlass::gemm::device::Gemm模板类来调用需要指定几个关键模板参数ElementA、ElementB、ElementC、LayoutA、LayoutB、LayoutC、ArchTag 等。对于 MXFP4ElementA 和 ElementB 是cutlass::mx_float4_t或者类似的类型缩放因子类型通过ScaleA、ScaleB指定。一个典型的调用流程是这样的先定义 ProblemShapeM、N、K然后配置 TileShape比如 128×128×64再选择 ClusterShape 和 StageCount。TileShape 的 K 维度必须是块大小的整数倍比如块大小 32那 K tile 至少是 32 或 64。StageCount 决定了流水线深度通常 3 到 4 比较合适太深会占满共享内存。下面是一个简化的代码框架展示 CUTLASS FP4 GEMM 的基本结构using ElementA cutlass::mx_float4_tcutlass::float_e2m1_t; using ElementB cutlass::mx_float4_tcutlass::float_e2m1_t; using ElementC float; using LayoutA cutlass::layout::RowMajor; using LayoutB cutlass::layout::ColumnMajor; using LayoutC cutlass::layout::RowMajor; using GemmKernel cutlass::gemm::device::Gemm ElementA, LayoutA, ElementB, LayoutB, ElementC, LayoutC, float, // Accumulator cutlass::arch::OpClassTensorOp, cutlass::arch::Sm100 // 或对应的 Rubin 架构标签 ; using GemmArguments typename GemmKernel::Arguments; GemmArguments args{ cutlass::gemm::GemmCoord{M, N, K}, tensor_a, tensor_b, tensor_c, tensor_d, scale_a, scale_b, alpha, beta }; GemmKernel gemm_op; auto status gemm_op.can_implement(args); if (status ! cutlass::Status::kSuccess) { /* 处理错误 */ } gemm_op.initialize(args); gemm_op.run();注意Sm100是 Blackwell 的架构标签Rubin 对应的标签需要等 CUTLASS 更新。在写代码之前先确认你用的 CUTLASS 版本是否支持目标架构否则编译会直接报错。4. 完整实操流程从数据准备到性能验证4.1 环境准备与依赖版本确认动手之前先把环境理清楚。你需要 CUDA Toolkit、CUTLASS、以及一块支持 FP4 的 GPU。CUDA 版本至少是 12.8 以上CUTLASS 建议用最新的 main 分支或者 3.x 的稳定版。如果你用的是 Blackwell 消费卡比如 RTX 50 系要注意一点消费卡的 FP4 支持可能和 datacenter 卡有差异某些 CUTLASS 的 FP4 kernel 在消费卡上可能跑不起来或者性能不如预期。我实测下来用nvidia-smi确认 GPU 型号和驱动版本用nvcc --version确认 CUDA 版本然后在 CUTLASS 的 CMake 配置里打开对应的架构选项。比如cmake -DCUTLASS_NVCC_ARCHS100 -DCUTLASS_ENABLE_TENSOR_CORE_MMAON ..这里的100对应 SM100Rubin 的架构号需要等官方公布后再调整。编译 example 的时候优先跑examples/目录下的 FP4 相关示例确认基础功能正常再改自己的代码。4.2 量化权重与激活值的实操步骤假设你有一个 FP16 的模型权重想转成 MXFP4 来做推理。步骤大致如下第一步确定量化粒度。是按 per-tensor 量化还是 per-block 量化MXFP4 默认是 per-block块大小 32。如果你的框架支持直接用框架的量化工具如果不支持就得自己写脚本。第二步计算每个块的缩放因子。对于每个 32 元素的块找到块内绝对值最大的元素然后计算缩放因子使得该元素映射到 FP4 能表示的最大值。E2M1 的最大值是 6所以缩放因子 max_abs / 6。然后把这个缩放因子转成 E8M0 格式取最接近的 2 的幂次。第三步用量化后的 FP4 值和缩放因子替换原始权重。注意 FP4 的打包方式两个 FP4 值存一个字节低 4 位存第一个高 4 位存第二个。这个打包顺序要和 CUTLASS 的预期一致否则计算结果会完全错乱。第四步激活值的量化。激活值是运行时产生的所以需要在 kernel 里动态量化。CUTLASS 的 FP4 GEMM 通常支持在加载 A 和 B 的时候做在线量化但这会引入额外开销。如果激活值分布比较稳定可以预先校准一个全局缩放因子减少运行时计算。4.3 性能验证与精度对比方法跑通之后怎么判断结果对不对、性能好不好精度方面用 FP16 的结果作为基准计算 FP4 结果的相对误差。对于大多数层相对误差在 1% 以内是可以接受的对于 attention 的 QK^T误差要求更严可能需要 0.1% 以内。如果误差超标先检查缩放因子计算是否正确再检查累加器是不是 FP32。性能方面用nsys或者ncu来 profile。关注几个指标SM 利用率、Tensor Core 利用率、显存带宽利用率。FP4 GEMM 的理想状态是 Tensor Core 利用率接近 100%显存带宽成为瓶颈。如果 Tensor Core 利用率很低可能是 tile 配置不合理或者缩放因子的加载拖慢了流水线。我一般会做一个简单的对比表把 FP16、FP8、FP4 三种精度的吞吐和精度放在一起看精度吞吐相对值显存占用相对值典型相对误差FP161.01.0基准FP82.00.5 0.5%FP44.00.251% - 3%这个表是粗略估计实际数值取决于模型和硬件。但趋势是明确的FP4 用精度换吞吐和显存适合对精度不那么敏感的场景比如推荐系统、部分 CV 任务、以及大模型推理里的一些非关键层。4.4 一个完整的 FP4 GEMM 调用示例下面给一个更完整的示例展示从 host 端准备数据到 kernel 启动的流程。假设 M4096N4096K4096块大小 32。// Host 端准备数据 int M 4096, N 4096, K 4096; int block_size 32; int num_blocks_k K / block_size; // A 矩阵M x KFP4 打包每字节两个元素 std::vectoruint8_t h_A(M * K / 2); // B 矩阵K x NFP4 打包 std::vectoruint8_t h_B(K * N / 2); // 缩放因子A 是 M x num_blocks_kB 是 num_blocks_k x N std::vectoruint8_t h_scale_A(M * num_blocks_k); std::vectoruint8_t h_scale_B(num_blocks_k * N); // 填充数据这里省略量化过程假设已经量化好 // ... // 拷贝到 device uint8_t *d_A, *d_B, *d_scale_A, *d_scale_B; float *d_C; cudaMalloc(d_A, h_A.size()); cudaMalloc(d_B, h_B.size()); cudaMalloc(d_scale_A, h_scale_A.size()); cudaMalloc(d_scale_B, h_scale_B.size()); cudaMalloc(d_C, M * N * sizeof(float)); cudaMemcpy(d_A, h_A.data(), h_A.size(), cudaMemcpyHostToDevice); cudaMemcpy(d_B, h_B.data(), h_B.size(), cudaMemcpyHostToDevice); cudaMemcpy(d_scale_A, h_scale_A.data(), h_scale_A.size(), cudaMemcpyHostToDevice); cudaMemcpy(d_scale_B, h_scale_B.data(), h_scale_B.size(), cudaMemcpyHostToDevice); // 构造 CUTLASS 参数并启动 // 参考上一节的 GemmArguments 构造这个示例省略了量化细节因为量化本身可以写一整篇文章。重点是让你看到数据布局和缩放因子的组织方式。实际使用时A 和 B 的缩放因子布局要和 CUTLASS 的预期完全匹配否则结果会错。5. 常见问题与排查技巧实录5.1 结果全错或部分错乱的排查路径FP4 GEMM 最容易出的问题就是结果不对。排查的时候按这个顺序来先检查数据打包顺序。FP4 两个值存一个字节到底是低 4 位在前还是高 4 位在前CUTLASS 的文档里有明确说明。如果你自己打包的数据和 CUTLASS 预期相反结果会完全乱掉。验证方法很简单构造一个已知的小矩阵手动算一遍和 kernel 输出对比。再检查缩放因子的布局。A 的缩放因子是按行还是按列B 的呢CUTLASS 的 MXFP4 通常要求 A 的缩放因子按 M×K/32 布局B 的按 K/32×N 布局。如果搞反了部分结果会错而不是全错。然后检查累加器类型。如果累加器设成了 FP16K 较大的时候误差会累积。改成 FP32 再试。最后检查 tile 配置。TileShape 的 K 维度必须是块大小的整数倍否则边界处理会出问题。比如块大小 32K tile 设成 48 就会出错。5.2 性能不达预期的常见原因性能问题通常比精度问题更难查。我遇到过几种典型情况第一种是缩放因子加载成为瓶颈。如果缩放因子的内存访问不连续每次加载都会产生额外的内存事务。解决办法是调整缩放因子的存储布局让它和数据的访问模式对齐。CUTLASS 里可以通过ScaleA和ScaleB的 Layout 参数来控制。第二种是 tile 太小。TileShape 设成 64×64×32 的时候每个 tile 的计算量不够流水线填充和排空的开销占比太高。建议至少 128×128×64如果共享内存够的话可以上到 256×128×64。第三种是 StageCount 不合适。StageCount 太小流水线会断流太大共享内存不够会限制 occupancy。一般 3 到 4 是甜点区具体要看你的 tile 大小和共享内存容量。第四种是消费卡上的 FP4 支持不完整。前面提过RTX 50 系虽然基于 Blackwell但某些 FP4 指令可能被阉割。如果你在消费卡上跑 datacenter 卡的 kernel性能可能只有预期的几分之一。这种情况只能换卡或者等驱动更新。5.3 精度不达标时的调整策略精度不达标先从量化策略入手。如果块大小 32 不够可以试试块大小 16 或 8。块越小缩放因子越精细精度越高但缩放因子的存储开销也越大。块大小 16 的存储开销是块大小 32 的两倍。如果块大小调整后还不够可以考虑混合精度。比如对精度敏感的层用 FP8其他层用 FP4。CUTLASS 支持在同一个 kernel 里混合不同的精度吗目前来看同一个 GEMM 里混合精度比较麻烦但可以在模型层面做分层量化不同层用不同的 kernel。还有一个技巧是校准缩放因子。不要直接用 max_abs 算缩放因子而是用百分位数比如 99.9% 分位数。这样可以避免个别离群值把缩放因子拉得太大导致其他值量化误差增加。这个技巧在激活值量化上特别有用。5.4 常见问题速查表问题现象可能原因排查方法解决思路结果全错数据打包顺序反了用小矩阵手动验证调整打包顺序部分结果错缩放因子布局不对检查 ScaleA/ScaleB 的 Layout按 CUTLASS 要求调整误差随 K 增大累加器精度不够检查 Accumulator 类型改用 FP32 累加性能只有预期一半Tile 太小或 StageCount 不当用 ncu 看 Tensor Core 利用率增大 tile调整 stage消费卡上跑不动架构支持不完整查 CUTLASS 的架构要求换 datacenter 卡或等更新边界结果错K tile 不是块大小整数倍检查 TileShape 的 K 维度调整为块大小整数倍6. 从 Blackwell 到 Rubin 的迁移注意事项6.1 架构标签与编译选项的变化从 Blackwell 迁移到 Rubin最直接的变化是架构标签。Blackwell 是 SM100Rubin 的标签需要等 NVIDIA 公布。在 CUTLASS 里架构标签决定了用哪套 MMA 指令和内存加载指令。如果你用 SM100 的标签编译然后在 Rubin 上跑可能能跑但性能不是最优反过来则可能直接编译失败。迁移的时候先把 CUTLASS 更新到支持 Rubin 的版本然后改 CMake 里的CUTLASS_NVCC_ARCHS。如果 CUTLASS 还没正式支持可以关注它的 GitHub 仓库通常新架构的支持会先在 develop 分支出现。6.2 指令集差异与 kernel 适配Rubin 的张量核心指令集相比 Blackwell 可能有变化。比如 MMA 的形状、操作数寄存器数量、缩放因子的处理方式等。这些变化对上层应用是透明的但如果你在写自定义 kernel就需要关注。一个实际的建议是尽量用 CUTLASS 的高层接口不要直接写 PTX 或 SASS。CUTLASS 会帮你处理指令集的差异你只需要改模板参数。如果你确实需要写底层代码那就得等 PTX ISA 文档更新确认 Rubin 的指令格式。6.3 性能调优参数的重新校准Blackwell 上调好的 tile 配置和 StageCount在 Rubin 上不一定最优。因为 Rubin 的共享内存容量、寄存器数量、Tensor Core 吞吐都可能变化。迁移之后建议重新跑一遍 autotuning或者手动试几组配置。我一般的做法是先固定一个保守的配置比如 128×128×64StageCount3确认功能正确然后逐步增大 tile 和 StageCount观察性能变化。每次只改一个参数记录吞吐和精度找到甜点区。7. 个人实操体会与后续扩展方向我在实际用 FP4 GEMM 的过程中最大的体会是精度和性能的平衡不是靠一个参数决定的而是量化策略、kernel 配置、模型结构三者共同作用的结果。同样的 FP4 kernel用在不同的层上效果可能天差地别。所以不要指望一套配置打天下该做的实验一个都省不了。另一个体会是CUTLASS 的文档虽然全但 FP4 相关的部分更新很快有时候文档和代码不一致。遇到这种情况直接看 example 的源码比看文档靠谱。example 里的配置是经过验证的照着改不容易出错。后续如果想深入可以往几个方向扩展。一是研究混合精度 GEMM在同一个 kernel 里对不同 K 块用不同精度。二是探索 FP4 在训练场景的可行性目前 FP4 主要用于推理训练还是 BF16 为主但未来不一定。三是关注 CUTLASS 对 Rubin 的正式支持等官方 example 出来之后第一时间跑一遍把踩坑经验补上。最后分享一个小技巧调试 FP4 GEMM 的时候先把 M、N、K 都设成很小的值比如 64然后用 CPU 写一个参考实现逐元素对比。这样能快速定位是数据布局问题还是计算逻辑问题。等小矩阵对了再放大到实际尺寸问题会少很多。