☰
Rubin架构FP4 GEMM深度解析:低精度推理如何成为GPU新赛道
2026/9/26 6:42:48 网站建设 项目流程

最近圈子里讨论最热闹的下一代 GPU,毫无疑问是 NVIDIA 的 Vera Rubin 平台。我翻了不少公开资料和技术拆解,发现大家对 Rubin 最关心的点,基本都落在 FP4 GEMM 上。FP4 这个精度,从 Blackwell 开始正式走上 Tensor Core 的舞台,到了 Rubin 这里,它不再是边角料,而是直接被当成性能指标的“扛把子”。那么 Rubin 的 FP4 GEMM 到底有什么特殊的?为什么它能成为下一代推理优化的核心赛道?如果你是做模型推理优化、量化部署,或者单纯对 GPU 架构感兴趣,这篇“概览”型的文章应该能帮你把零散信息串起来。先说明一下,本文不是官方文档,很多地方是结合公开信息和架构演进的合理推断,我会把这些推断明确标出来,大家辩证着看。

1. 为什么 FP4 GEMM 成了 Rubin 的代名词

1.1 先把 FP4 和 GEMM 这两个黑话掰开

FP4 就是 4 位浮点数,通常采用 E2M1 格式:1 位符号、2 位指数、1 位尾数。相比于 FP8(E4M3 或 E5M2)和 BF16,FP4 的位宽只有它们的二分之一或四分之一,能表示的数据范围和精度自然更有限。GEMM 是 General Matrix Multiplication,也就是通用矩阵乘。在深度学习里,线性层、卷积、注意力机制中的投影和打分操作,本质上都能归结成矩阵乘。所谓 FP4 GEMM,就是让两个 4 位浮点矩阵直接相乘,而不是像以前那样先把数据转回 FP16 再算。

但为什么 FP4 和 GEMM 要放在一起说?因为神经网络里绝大多数计算量都集中在 GEMM。把 GEMM 的输入精度砍到 4 位,等于让同一个 Tensor Core 时钟周期内能处理的元素数量翻倍甚至翻四倍。按位宽估算,FP4 乘法器的面积、功耗都比 FP16 小得多,因此单位芯片面积上可以塞进更多的计算单元。Rubin 把 FP4 GEMM 作为宣传亮点,意味着它在硬件层面不只是“支持一下”,而是专门为 4 位矩阵乘优化了微架构。

这里还要澄清一个常见误区:FP4 GEMM 不等于“结果也存成 FP4”。矩阵输入虽然是 4 位,但内部累加器通常保留 FP32 精度的中间结果,最后再通过缩放因子和舍入转换才写回低精度。也就是说,FP4 GEMM 是一个“输入低精度、输出高精度累加”的过程。这个设计思路从 FP16 时代一直延续到现在,只是输入位宽一步步从 16 降到 8 再降到 4。

1.2 量化的大趋势:AI 推理正在被低位宽重塑

现代大模型动辄几百亿参数,推理时最大的痛点不是“算不动”,而是“显存装不下、带宽跟不上”。GPT 这类自回归模型在 decode 阶段,一个 token 一个 token 地生成,每次都要把整个模型权重从显存搬到计算单元。权重位宽减一半,带宽占用就减一半;缓存容量也能塞下更多权重。所以近几年 INT8 量化、INT4 量化早已是工程标配。

FP4 相比 INT4 有一个明显的优势:指数位带来了动态范围。权重和激活值经常会出现少量特别大或特别小的离群值,纯整数的 INT4 遇到这些值很容易直接溢出或饱和,浮点格式则能通过指数把数值范围拉得更宽。虽然 FP4 的精度仍然很粗糙,但在配合缩放因子的情况下,经验上很多模型都能承受这种损失,尤其是加入量化感知训练后,精度可能逼近 FP16 基线。

这也就解释了为什么 Blackwell 开始引入 FP4 Tensor Core,Rubin 又进一步把它做强。推理侧的性价比模型已经从“能用 FP16 就行”变成了“能上 FP4 就上 FP4”。当大家都把模型量化到 4 位时,GEMM 的性能就直接决定了整个推理系统的吞吐上限。

1.3 算力参数表的暗示:FP8 是训练主菜,FP4 是推理甜品

NVIDIA 历代 GPU 的规格表里,峰值算力通常按 FP32、FP16、FP8、FP4 等分列。Hopper H100 只提 FP8 Tensor Core,Blackwell B200 把 FP4 加进了指标中,而 Rubin 的官方预告里,FP4 被提到了一个非常显眼的位置。这不是简单的市场宣传,而是明确告诉你:Tensor Core 的物理设计已经预留了 FP4 的原生路径。

我个人的理解是,FP8 依然会是训练阶段的主力,因为反向传播对梯度精度要求更高,FP4 很难撑起大批量训练。但在推理阶段,模型权重已经固定,可以通过离线量化把权重压到 FP4,激活值也能节省大量带宽。Rubin 的定位很可能是“推理专用算力大幅提升,同时保持 FP8 训练能力”。如果你看到未来发布规格中 FP4 稀疏算力是 FP8 的两倍以上,完全不用惊讶,这正是架构层面的有意取舍。

2. Rubin 架构怎么为 FP4 GEMM 铺路

2.1 Vera Rubin 平台概览

Vera Rubin 是 NVIDIA 下一代计算平台的名字,包含 Vera CPU 和 Rubin GPU 两部分。Rubin GPU 采用双裸片(chiplet)设计,两个计算 die 之间通过 NV-HBI 高带宽接口互联,搭配 HBM4 显存。整个平台还引入了 NVLink Hub,用于多 GPU 之间的互联拓扑。这些听起来和 FP4 GEMM 没直接关系,但实际上环环相扣。

GPU 从一个架构走到另一个架构,最大瓶颈往往是“喂给计算单元的数据不够快”。Rubin 之所以把 HBM 从 HBM3E 升到 HBM4,把互联从传统 PCB 走线改成更先进的基板封装,就是为了让每个时钟周期都能搬运足够多的低精度数据。FP4 GEMM 的计算密度异常高,如果显存带宽不够,芯片就会大量空转,浪费时间在等数据上。

这里值得提一句:Vera Rubin 的命名来自天文学家 Vera Rubin,这也暗示 NVIDIA 想强调“探索未知”的调性。对我们做工程的人来说,名字不重要,重要的是它背后的物理设计是否能真正兑现超过上一代的推理性能。

2.2 Tensor Core:FP4 乘法器密度怎么玩

Tensor Core 是目前 GPU 上最适合矩阵乘的专用单元。从 Volta 时代起,它就固定按小 tile 进行矩阵乘。一般来说,一个 Tensor Core 内部由多个 4x4 或 8x8 的乘累加阵列组成,每个周期算出一个子矩阵块的部分结果。Hopper 的 Tensor Core 针对 FP8 做了专门设计,Blackwell 则扩展到 FP4。

对 FP4 来说,一个乘法器只需要处理 4 位尾数(实际上尾数只有 1 位),芯片上可以放更小尺寸的乘法阵列。假设上一个架构用 4 位乘法器做 FP16,需要 16 位乘法器执行两次 4 位乘法?不,合适的方式是:FP4 位宽减半,单次乘法需要的晶体管数量更少,因此相同面积下能塞入更多并行单元。Rubin 大概率会针对 FP4 增加独立的张量核心通路,或者让同一 Tensor Core 在 FP4 模式下按更大的 K 维度展开计算。

这里有一个合理推断:当前的 Tensor Core 指令通常以 m16n8k8 之类的 tile 为单位,其中 k 是缩减维度。对于 16x8x8 的运算,如果每个数都是 4 位,数据位宽大约是 FP16 的四分之一,那么内部计算阵列有能力把 k 扩展到 32 甚至 64,一次性算完更长的规约链。这会让 FP4 GEMM 的指令级并行度进一步提升,但也会给寄存器堆和累加器带来压力。

2.3 带宽、显存与 PCB:FP4 GEMM 的运行环境

讲 FP4 GEMM 时,很多人只盯着“算得快”,但真实瓶颈经常是数据搬运。假设某个线性层的权重矩阵是 10000x10000,FP16 下需要约 200MB 数据,FP4 下只需要约 50MB,但硬件吞吐翻倍后,单位时间内消耗的数据量可能是原来的两倍多。也就是说,FP4 不仅没有轻松,反而更依赖显存带宽。

所以 Rubin 必须搭配 HBM4、更宽的缓存和更高效的片上互连。热搜词里有“英伟达 rubin pcb”,这也是非常关键的一点。PCB 和基板的设计直接影响信号完整性、电源完整性和散热能力。Chiplet 之间通过基板上的高速走线进行通信,如果布线长度不匹配、串扰控制不好,信号频率就提不上去。电源网络如果去耦电容不够,瞬间电流波动会让 GPU 降频。换句话说,FP4 的原始算力再高,供电送不进去或者信号乱了,实际表现也出不来。

PCB 不是把芯片装起来那么简单,它本质上是一个高速信号传输系统。Rubin 这类 2000W 级别功耗的 GPU,PCB 上的电源层、地层设计、过孔位置、阻抗匹配都必须按许多 GHz 甚至几十 GHz 的标准去做。这是从 Blackwell 到 Rubin 一直都在强攻的硬骨头,只是普通用户很难直接感知罢了。

3. FP4 GEMM 的硬件实现思路

3.1 FP4 数值格式:为什么必须引入缩放因子

FP4 的 E2M1 格式只有 1 位尾数,符号 1 位,指数 2 位,指数偏置为 1。它可以表示的有限值数量非常少,最大有限值约等于 6.0,最小正常值约 0.5。如果待计算的数据分布在 0.001 到 1000 这个范围,直接存成 FP4 几乎一半数据会变成 0 或 6.0,损失惨重。

解决方式是给数据加一个缩放因子。常见做法是按照 block 为单位(比如 32 个元素一组),为这一组元素计算一个公共的缩放因子,然后用 FP4 存储规约后的数值。这个思想来自 OCP 制定的 Microscaling(微缩放,MX)格式,典型的有 MXFP4、MXFP8。缩放因子通常用 E8M0 这类纯指数格式存储,只记录 2 的幂次,这样硬件只需要做指数加减,成本可以压得非常低。

在 GEMM 当中,A 矩阵和 B 矩阵各自有缩放因子,相乘之后整体结果还要乘以两个缩放因子的乘积。因此 Tensor Core 的管线里不只是简单乘加,而是要在累加结果上执行一个“缩放重加载”操作。这个操作可以放在 epilogue 阶段,也可以提前在数据进入 Tensor Core 之前处理。无论哪种方式,缩放因子的优化设计都决定了 FP4 GEMM 能带来多少收益。

3.2 GEMM 流水线:从加载到累加的完整路径

一个 FP4 GEMM 在 GPU 上的执行过程,大致可以分成以下几步:全局内存读取 A、B 和各自的缩放因子,写入共享内存;从共享内存按 tile 读取,经寄存器缓存后压入 Tensor Core;Tensor Core 执行矩阵乘和累加,得到 FP32 中间结果;epilogue 阶段把 FP32 结果乘以缩放因子、加 bias、过激活函数,最后按需转成 FP4 输出。

与传统 FP16 GEMM 最大的不同在于,FP4 经过的每个数据通路都要“窄一点”。例如,从显存到共享内存的传输中,一个 32 字节的请求可能携带 64 个 4 位元素。在寄存器阶段,打包和解包 FP4 数据会成为额外开销。如果硬件不能高效处理这种半字节对齐,软件层就需要自己做位压缩,反而拖慢性能。因此 Rubin 的 Tensor Core 很可能在 Load 阶段就支持从内存读取紧凑的 4 位数据,并在寄存器中自动解包成内部更宽的表示。

另一种思路是把 FP4 直接映射到可寻址的 FP8 或 FP16 通道中,比如一个 FP4 元素占半个 FP8 槽,两个 FP4 拼成一个 FP8。这样硬件可以复用已有的低精度数据通路,通过多路选择器和移位寄存器把两个 4 位值并行送入乘法器。这类实现属于微架构细节,我们看不到,但从软件性能参数上可以推测,Rubin 的 FP4 峰值算力至少是 FP8 的两倍,大概就是走了这类“位宽复用”的路线。

3.3 2:4 稀疏 + FP4:稀疏 GEMM 的乘法吞吐

NVIDIA 从 Ampere 架构开始支持 2:4 结构性稀疏:每 4 个元素中只保留 2 个非零值,且零值的位置固定。稀疏 GEMM 可以跳过一半的乘法运算,理论上吞吐能翻倍。对 FP4 来说,这种稀疏性同样重要,因为 4 位权重本来就很小,剪掉一半的权重大幅减少计算量。

但稀疏性也有代价。稀疏矩阵必须额外存储索引或位掩码(metadata),在加载时需要读取额外数据。如果带宽本来就捉襟见肘,稀疏化省下的计算时间可能被额外 metadata 的带宽开销抵消。实际中,2:4 稀疏更适合权重矩阵,因为权重可以提前离线剪枝;而对于激活值,由于每个 batch 不同,很难固定零值位置,所以很少用结构化稀疏。

Rubin 大概率会延续 Blackwell 上的稀疏 Tensor Core 方案,同时支持 FP4 和稀疏性混用。官方规格里的“峰值稀疏 FP4 TFLOPS”通常比稠密 FP4 高一倍,但这个数字是理论极限,真实场景能不能跑到一半以上,要看数据分布和软件是否能把稀疏 pattern 高效映射到硬件。

3.4 混合同精度 GEMM:权重和激活精度不对称

在实际部署中,权重和激活值的敏感度往往不一样。有些层权重很稳定,压到 FP4 几乎没有问题,但激活值波动大,用 FP4 会掉点,需要保留 FP8。那么硬件能不能让 A 矩阵用 FP4,B 矩阵用 FP8?理论上 GEMM 的硬件乘法器应该能支持这种混合输入,只需要把低精度的数据提升到同一个中间精度,然后相乘。

Hopper 的 FP8 模式就允许 A 和 B 使用不同的 E4M3/E5M2 格式。Blackwell 在 FP4 上很可能会扩展出混合精度组合,比如 A:FP4, B:FP8,或者 A:FP8, B:FP4。这样一来,软件栈就可以按层选择最优的量化方案,而不是一刀切全用 FP4。Rubin 的指令集如果能提供这种灵活性,对精度和性能的平衡会带来很大帮助。不过这也意味着 Tensor Core 的内部乘法器需要支持不同位宽的输入通道,硬件复杂度会上升。

4. 软件栈与量化系统的配合

4.1 从一行代码到 Tensor Core,要经过多少层

用户写一句torch.mm(a, b),看起来很简单,但底层会经历 PyTorch 算子分派、ATen 矩阵乘选择、cuBLAS/cuBLASLt 库调用、CUDA 驱动,最后才能落在 SASS 指令上。FP4 GEMM 的调用链比普通 FP16 更长,因为中间还要传递量化参数、缩放因子和存储格式。

如果应用程序直接调用 cuBLASLt,需要配置一个matmulDescriptor,指定数据类型为 FP4,并提供缩放因子的维度信息。框架层的 TensorRT 或 vLLM 也都在为这种新精度做适配。就我目前看到的生态进展,FP4 推理还没有达到 FP8 那样的成熟度,很多量化工具都还在 preview 阶段。Rubin 真正发布后,CUDA 工具包、cuBLAS、CUTLASS 都会快速跟进,但在早期阶段,工程师可能得自己写一些自定义 kernel 或者 workaround。

关键建议是:不要等硬件到了才开始调软件。现在就可以在模拟器或现有硬件上做 FP4 数值仿真,确定哪些层适合 FP4、哪些层需要保留 FP8。软件栈成熟后,迁移成本会低很多。

4.2 缩放因子放哪里:架构层面的 epilogue 设计

FP4 GEMM 的输出节点通常要考虑 scale 和 bias 的融合。比如输出 C = (A_fp4 * B_fp4) * (scaleA * scaleB) + bias,这个计算如果在 GEMM 内部完成,可以避免多次访问显存。cuBLASLt 的 epilogue 设计支持不同组合:不融、加 bias、加 scale、加激活函数等。

在实际工程里,我建议尽量把缩放因子的“数据类型”设计成 E8M0 这种纯指数整数,因为它和任何尾数无关,乘法就是指数相加,非常简单。如果缩放因子也用 FP8 的 E4M3 格式,硬件还得处理尾数乘法,反而会引入额外误差和延迟。这不是一定的,但 OCP MX 规范里 E8M0 就是设计来干这件事的。

另一个细节是缩放因子的粒度。按 tensor 缩放最简单但精度最差,按 channel(也称为 per-channel)压缩权重的精度更好,按 block 缩放最强但需要额外存储。PyTorch 中常见的 weight_quant 实现就支持 per-group scale。在 FP4 GEMM 里,缩放因子本身也是要参与计算的,所以它的存储布局会影响 GEMM kernel 的读取效率。最理想的情况是缩放因子随矩阵一起按 tile 加载,避免二次访存。

4.3 QAT 训练模拟:让 FP4 模型不掉点

如果你只是把训练好的 FP16 模型直接转成 FP4,大概率会掉点,尤其是激活值比较大的任务。最好的做法是使用量化感知训练(QAT),在训练前向过程中用fake_quantize模拟 FP4 的舍入误差,让模型自己去适应。

一个常见误区是:QAT 模拟时直接把 FP4 的结果当浮点用,却忽略了缩放因子的动态更新。你需要用直通估计器(STE)让梯度绕过量化节点,并且缩放因子要根据当前激活的统计量实时调整。像quantization.observer里的 MinMaxObserver、PerChannelMinMaxObserver 等等,都可以用来估计缩放。但要注意,FP4 的动态范围太窄,单纯统计 min/max 可能会被离群值带偏。更稳妥的做法是使用百分位剪裁,比如 99.9% 分位点作为缩放上限,这样既保留大部分数据,又容忍极端值溢出。

QAT 训练中如果硬要精确模拟 FP4 GEMM 的硬件行为,比较麻烦,因为 Tensor Core 内部的舍入可能和软件模拟不一样。我通常建议先按 E2M1+缩放因子的公式写一个并行 CPU/GPU reference,再和预期结果对照。这个参考实现的性能不需要多好,但正确性必须高。

4.4 伪代码:如何配置一个 FP4 GEMM 调用

假设我们要在 cuBLASLt 里执行一个 FP4 GEMM,伪代码逻辑如下:

// 这里以 cuBLASLt 风格示意,实际 API 以最新版为准 cublasLtMatmulDesc_t opDesc; cublasLtMatmulDescCreate(&opDesc, CUDA_R_32F, CUDA_R_32F); // 设置 A、B 均为 FP4,输出为 FP32 cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_A_TYPE, &CUDA_R_4F, sizeof(CUDA_R_4F)); cublasLtMatmulDescSetAttribute(opDesc, CUBLASLT_MATMUL_DESC_B_TYPE, &CUDA_R_4F, sizeof(CUDA_R_4F)); // 指向缩放因子 cublasLtMatrixLayout_t scaleLayout; // 设置 scale 的维度、stride,指向 scaleA 和 scaleB

关键点是scaleLayout必须和矩阵的布局对齐。如果缩放是 per-channel 的,那么 scaleA 的大小就是 M 或 K 的长度;如果 per-block,需要额外传入 block 尺寸。API 调对了,底层 kernel 才能在所有 tile 上高效读取 scale。

真实项目里,我反而建议先用 TensorRT 的 Q/DQ 节点或 vLLM 的 FP4 实现,进一步向下调封装,那会让开发效率高很多。自己写 cuBLASLt 调用虽然灵活,但要处理太多格式细节,比如 FP4 数据的位序、端序、对齐,稍不留神就出 bug。

5. 工程实践中的性能调优与避坑

5.1 先判断是带宽受限还是算力受限

每次做 FP4 GEMM 优化,我第一件事就是算计算强度(arithmetic intensity)。计算公式是:计算强度 = FLOPs / Bytes。如果这个值小于机器的“机械特性强度”,那这个 GEMM 就是带宽受限;反之则是算力受限。

举个例子,一个 4096 x 4096 的矩阵乘,FLOPs 约等于2*4096^3 = 1.37e11。如果 A 和 B 都是 FP4,每个元素只有 0.5 字节,总共需要读取2*4096*4096*0.5 = 16.8MB。那么计算强度大约是 8155 FLOPs/byte。再看机器算力:假设峰值 FP4 是 50 TFLOPS,带宽是 8 TB/s,那么机械特性强度是 6250 FLOPs/byte。这个例子中计算强度略高于机械特性,它在理论上偏向算力受限。但一旦你引入 padding、metadata、非对齐访存,实际有效带宽会打折,很有可能变成带宽受限。

所以优化方向在不同卡上完全不一样。带宽受限时,优先做数据压缩、减少 metadata、使用 TMA 异步复制;算力受限时,优先调整 tile 大小、提高 Tensor Core 利用率、减少寄存器浪费。不能一把梭。

5.2 半字节存储和内存对齐的麻烦

FP4 数据是 4 位,两个 FP4 放在同一个字节里。如果你按常规的uint8_t数组存储,低位是第一个元素还是高位是第一个元素,在不同 API 里可能有不同约定。CUTLASS 通常建议用uint8_t表示两个元素,但具体 bit 顺序需要自己核对。

更麻烦的是矩阵的 K 维或 M 维如果不对齐到 2 的倍数,最后一个字节会有一半是空数据。你必须在分配内存时做 padding,也就是把 stride 适当加大。我习惯把所有维度都对齐到 16 或 32 的倍数,虽然会浪费一点显存,但能避免很多 kernel 边界问题。

另外一个坑是共享内存的 bank conflict。FP4 数据在共享内存中打包得非常紧凑,读取时如果多个线程访问同一个字节的不同半字节,可能会产生 bank conflict。最佳实践是把打包后的数据先按 uint8_t 数组加载到寄存器,再在寄存器内部做位移和掩码,而不是让共享内存硬件去猜你要哪个半字节。这样可控性最高,性能也更稳定。

5.3 稀疏性不是白拿的:metadata 也会占带宽

2:4 结构稀疏的 FP4 GEMM,理论上乘法任务少一半,但代价是要额外读取 metadata。每个 4 元素组用 4 位或 8 位保存哪些位置是非零值,所以稀疏矩阵的存储从原来的 2 nibbles 变成了额外 metadata。如果稀疏矩阵本身非零值比例高,或者零值分布不是 2:4 结构,那就要重新压缩,否则收益基本为零。

我的建议是在模型权重上做一次全局稀疏率统计,如果某层剪枝后只有 30% 的零值,那硬套 2:4 会打乱数据,带来精度损失,但性能收益不明显。对于权重稀疏率超过 60% 的层,再开 2:4 稀疏。对于激活值或 KV Cache,除非专门做过结构化剪枝,否则别碰稀疏。

5.4 功耗、散热与供电对峰值性能的影响

FP4 GEMM 的高吞吐会让 GPU 瞬时功耗明显上升,尤其是跑那种大型密集矩阵乘时,芯片内部数百个 Tensor Core 同时翻转,电流冲击非常剧烈。Rubin 如果 PCB 上的电源稳压模块、去耦电容、陶瓷电容布局不合理,电压纹波会比较大,GPU 只能靠降频来保护硬件。

我在之前做 Blackwell 性能测试时,把功耗墙从 100% 调到 80%,FP4 峰值算力骤降了 15% 以上,而且 FP8 的下降幅度没那么明显。原因是 FP4 GEMM 的并行度更高、功耗波动更剧烈,供电跟不上时 Boost 频率很难维持。因此评估 Rubin 的 FP4 性能,一定要看持续功耗下的实际吞吐,不要只看规格表峰值。如果你有整机柜部署计划,电源和散热预算也要按 FP4 满负荷去设计,留足冗余。

6. FP4 GEMM 的场景落地与下一步

6.1 场景图谱:哪些任务真正需要 FP4

我做了一个简单的适用性表格,可以参考:

场景FP4 GEMM 适用性原因
LLM 推理(大量并行请求)高权重和 KV Cache 都适合低精度压缩,吞吐收益大
图像生成(Stable Diffusion 类)中高UNet 和 Transformer 块的线性层很多,但激活值波动大,可能需要混合精度
推荐系统(大规模 embedding + MLP)高embedding 严重带宽受限,FP4 能大幅减少内存占用
科学计算低精度需求高,FP4 太少
自动驾驶边缘推理低延迟敏感,FP4 对异常天气场景的鲁棒性不足

这里有一个规律:凡是“数据量大、算力相对充裕、精度容忍度较高”的任务都适合 FP4 GEMM。典型的就是大模型推理和推荐系统,它们的高吞吐靠的是大批量把内存里的数据灌给计算单元,而不是在这几个比特里抠精度。

6.2 LLM 推理的数据流建议:Prefill 用 FP8、Decode 用 FP4

在 LLM 推理中,Prefill 阶段要处理很长的 prompt,计算强度高,适合用 FP8 甚至 FP16 保持较高精度。Decode 阶段每次只生成一个 token,主要从显存捞权重和 KV Cache,带宽占用远大于计算需求,这时候用 FP4 量化权重和 KV Cache 能显著降低延迟。

Rubin 的软件栈如果要动态在 FP8 和 FP4 之间切换,那会非常有意思。量化切换本身有开销,但如果模型已经按层做了 QAT,提前确定好哪些层用 FP4、哪些层用 FP8,运行时就不需要重新量化,只需要切换 GEMM kernel 即可。这也是批量推理引擎将来最需要的功能之一。

6.3 对未来训练和微调的影响

FP4 GEMM 直接用稳定训练目前还比较困难。反向传播需要计算梯度和更新权重,梯度范数变化很大,4 位浮点的动态范围不够。不过有一些研究方向在尝试“低精度梯度压缩”,用 FP4 近似梯度,再配合误差反馈补偿,这样也许能在分布式训练中减少通信量。

微调场景下,可以采用“冻结权重为 FP4 表示,但保留一份 FP32 的影子权重”的做法。前向推理用 FP4 GEMM,反向传播时用影子权重的高精度梯度。这本质上是一种混合精度微调,显存占用虽然增加了,但能让大模型在有限的 GPU 上做轻量适配。Rubin 的高带宽架构对这种“影子权重”模式很友好,因为后面更新梯度需要读取 FP32 权重,带宽如果不够反而会拖慢训练。

结尾:一点个人体会

最后说点个人经验。我在做量化部署时,第一眼看到 FP4 觉得真是好东西,但实际踩坑之后发现,FP4 GEMM 的成功不是靠硬件一个维度,而是数值格式、软件栈、供电和散热整体配合的结果。我的建议是,在 Rubin 真正大规模出货之前,先把你的模型在 FP4 模拟器上跑一遍,记录误差曲线和内存分布。不要等硬件到了才开始调内核,那时候就晚了。

另外一个实用技巧是,如果你打算用 FP4 GEMM,尽量把缩放因子设计成 8 位指数(E8M0),这样硬件可以用极低成本完成缩放,还能减少数值溢出的概率。后面我也会持续跟进 Rubin 的相关资料,有新发现再和大家分享。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询