CUDA 算子工程:手写 FlashAttention v2 之路
第 9 章 量化 Kernel:INT8 / FP8 / INT4
量化省的不只是显存,更是带宽——而带宽往往是 LLM 单 token 解码里最先用完的东西。 这一章引用的 Marlin 源码都可以在 vLLM v0.8.5 的
csrc/quantization/下逐行对上。
9.1 为什么 LLM 推理必须量化
来看一组数字。LLaMA-2 70B 模型,FP16 权重大小:
70B 参数 × 2 字节/参数 = 140 GB
H100 一张卡只有 80 GB HBM——70B 模型根本放不进单卡。如果 INT8 量化:
70B × 1 字节 = 70 GB 可以放进 H100 80GB 卡!
INT4 量化:
70B × 0.5 字节 = 35 GB 单张 H100 绰绰有余, 两张 24GB 消费卡也放得下
但部署成本还不是量化的最大动机——它真正解决的是带宽瓶颈。
LLM 单 token decoding 的本质是"读一遍权重做一次 GEMV"。一个 70B 模型每 token 要读 140 GB 权重(FP16)/ 70 GB(INT8)/ 35 GB(INT4)。在 H100 SXM5 3.35 TB/s 的 HBM3 上:
| 量化 | 每 token HBM 流量 | 纯带宽下限 | 加上开销后的量级(非实测) |
|---|---|---|---|
| FP16 | 140 GB | ~42 ms | ~50-60 ms |
| INT8 | 70 GB | ~21 ms | ~25-30 ms |
| INT4 | 35 GB | ~10 ms | ~12-15 ms |
INT4 的每 token HBM 流量只有 FP16 的 1/4,带宽 bound 的解码延迟也随之接近 1/4——这就是量化在 LLM 推理上的现金价值。
9.2 量化基础:把 FP16 压成整数
量化的核心数学:
其中:
- = scale(浮点缩放因子)
- = zero-point(整数零点偏移)。注意符号约定:这里用的是 GPTQ / AWQ / Marlin 代码里的通行写法——量化时加 、反量化时减 ,所以后面 9.6 节那段 Marlin 代码的注释才说它"把 这个对称零点融进了转换"。有的论文写成量化减、反量化加,两套式子各自自洽,混用才会出错。
- INT8: 或无符号
- INT4: 或无符号
例子:把 FP16 范围 [-2.0, +2.0] 压到 INT8:
fp16 = 1.5 -> int8 = round(1.5 / 0.0157) = 96
fp16 = -0.5 -> int8 = round(-0.5 / 0.0157) = -32
dequant:
int8 = 96 -> fp16 = 0.0157 * 96 = 1.507 (误差 0.007)
这就是量化的"信息损失"——0.007 的误差。但只要总体上模型仍能给出合理输出,就值得。
9.3 量化方案的"粒度"
scale 怎么选?这决定了量化方案的"粒度",权衡精度和元数据开销:
flowchart TB
subgraph Per-tensor [Per-tensor 一个 scale]
PT1[整个 weight 矩阵共享一个 s]
PT2[元数据极少, 1 个 fp32]
PT3[精度差, 大异常值会主导 scale]
end
subgraph Per-channel [Per-channel 每行/列一个 scale]
PC1[每个输出通道一个 s]
PC2[元数据 N_out × 4 字节]
PC3[精度好, 工业标配]
end
subgraph Per-group [Per-group 每 K 列一组 scale]
PG1[每 128 列一组 s]
PG2[元数据 N × N_groups × 2 字节, FP16 scale]
PG3[精度更好, INT4 必备]
end
subgraph Per-token [Per-token 推理时每行一个 s]
PTK1[激活每个 token 一行一个 s]
PTK2[运行时计算 s]
PTK3[配合 SmoothQuant]
end
不同方案的精度-开销权衡:
| 方案 | 精度损失 | 元数据 | 适用 |
|---|---|---|---|
| Per-tensor | 高 | 极小 | 老的部署,不推荐 |
| Per-channel | 中 | 小 | INT8 PTQ 标配 |
| Per-group (g=128) | 低 | 中 | AWQ / GPTQ INT4 标配 |
| Per-token | 低 | 运行时算 | SmoothQuant W8A8 |
INT4 通常要用 per-group——因为 INT4 范围太窄(只有 16 个值),单个 channel 内不同列的数值范围差异会让 per-channel 的精度明显下降。AWQ 论文全程用 group_size=128,GPTQ 模型的常见发布配置也是 128,意思是沿输入维每 128 个连续元素共享一个 scale。
9.4 LLM 量化算法概要
LLM 量化论文很多,工业上最常用的几个:
9.4.1 SmoothQuant (W8A8)
权重和激活都量化到 INT8。核心 insight:激活的异常值(outlier)远大于权重,但通过把"激活的难度"挪一部分到"权重的难度",可以让两边都好量化。
具体做法:用一个对角矩阵 把激活除回去、权重乘进去:
数学上完全等价,但 的范围变小了(容易 INT8), 的范围只略微变大(仍能 INT8)。
9.4.2 GPTQ (W4A16)
权重 INT4,激活 FP16。GPTQ 按固定顺序逐列量化——量化完第 列后,根据 Hessian(逆)信息更新后面尚未量化的列,补偿这一列引入的误差(论文特意放弃了前身 OBQ 的贪心顺序,证明任意固定顺序效果相近)。开源工具 AutoGPTQ(auto_gptq)是常用实现之一。
9.4.3 AWQ (W4A16)
Activation-aware Weight Quantization(Lin et al., 2023)。核心观察是不是所有权重都同样重要——按激活幅度能挑出约 1% 的"显著权重",把它们留在 FP16 就能大幅挽回精度。但论文明确指出这种混合精度数据类型会让系统实现变得困难,AWQ 实际采用的不是混合精度,而是等价的做法:给显著通道乘一个 per-channel scale(激活侧除回去),让它们在 INT4 网格里占到更多有效位。整个流程不需要反向传播,比 GPTQ 更简单,效果相当或更好。
9.4.4 FP8 量化
Hopper 引入的 FP8(E4M3/E5M2)是另一种思路——不再做整数量化,而是用更窄的浮点格式。FP8 的指数位提供更大的动态范围,对 LLM 里常见的带离群值分布通常比 INT8 更友好,但需要硬件原生支持(Ada 的 sm_89 与 Hopper 的 sm_90 起)。NVIDIA 的 Transformer Engine 和 H100 的 FP8 训练栈是这条路线。
9.5 关键 Kernel: Dequant + GEMM
量化 LLM 推理的核心 kernel 是 dequant-fused GEMM:
# 输入: 量化的 W (INT4 或 INT8) 和 FP16 激活 X
# 输出: FP16 结果 Y = X @ dequant(W)
# 朴素拆分:
W_fp16 = dequantize(W, scale, zero) # 临时分配 N×K 的 FP16
Y = X @ W_fp16 # cuBLAS GEMM
# Dequant-fused:
Y = quantized_gemm(X, W_int4, scale, zero)
朴素版本的问题:
- 临时分配 N×K 的 FP16 矩阵:以 LLaMA-2 70B 为例,8192×8192 的 attention 输出投影临时矩阵就有 128 MB(MLP 层是 8192×28672,达 448 MB),HBM 写一次。
- W_fp16 写完立刻读:完全是浪费 HBM 带宽。
- 失去了量化的带宽优势:cuBLAS GEMM 还是按 FP16 读 W_fp16,HBM 流量没省。
Dequant-fused 的核心思想:在 W 从 SMEM 取进寄存器、喂给 Tensor Core 之前才 dequant——HBM 和 SMEM 里放的都仍是 INT4,HBM 只读了 INT4 的流量(Marlin 就是这样:sh_b 里是打包的 INT4,dequant 在寄存器里做)。
flowchart LR HBM[HBM] -->|读 INT4 权重 + FP16 scale| SMEM[SMEM/Register] SMEM -->|dequant 在寄存器中| FP16[FP16 fragment] FP16 -->|喂给 Tensor Core| TC[mma.sync FP16] TC -->|FP32 accumulator| OUT[输出 D]
这是工业级量化 GEMM 的标准结构。Marlin(vLLM 里的算子名是 gptq_marlin_gemm)、Machete、TensorRT-LLM 的 weight-only GEMM 都是这种思路。
9.6 INT4 解包的黑魔法
INT4 dequant 中最关键的是怎么高效地从 packed INT4 解出 FP16。
INT4 用 4 位存储,每字节装 2 个 INT4。从 8 个 packed INT4 字节解出 16 个 FP16 是个"小问题"——但在 GEMM 内层这一步可能跑几亿次,必须做到极致。
朴素写法:
// 8 字节 packed INT4 -> 16 个 half
half decoded[16];
#pragma unroll
for (int i = 0; i < 8; ++i) {
int8_t packed = packed_data[i];
int8_t lo = packed & 0x0F; // 低 4 位
int8_t hi = (packed >> 4) & 0x0F;
decoded[i * 2 + 0] = __int2half_rn(lo - 8); // 中心化到 [-8, 7]
decoded[i * 2 + 1] = __int2half_rn(hi - 8);
}
这段代码 16 次类型转换,编译出来很多指令(nvcc 13.4 以 sm_90a 编译核对:16 条 I2F.F16,外加 24 条 LOP3.LUT 与 16 条 PRMT;不同 CUDA 版本可能不同)。性能差。
工业级写法用的是一个位运算技巧。下面摘自 vLLM v0.8.5 vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:169 起的 dequant<half, kU4B8>(lop3 本身在同文件 :138);代码逐字保留,行尾中文注释为本专栏所加,lop3 的 asm 约束合并成了一行:
template <int lut>
__device__ inline int lop3(int a, int b, int c) {
int res;
asm volatile("lop3.b32 %0, %1, %2, %3, %4;\n"
: "=r"(res) : "r"(a), "r"(b), "r"(c), "n"(lut));
return res;
}
// 从一个 int32 里的 4 个 nibble 解出 4 个 half
template <>
__device__ inline typename ScalarType<half>::FragB
dequant<half, vllm::kU4B8.id()>(int q) {
const int LO = 0x000f000f;
const int HI = 0x00f000f0;
const int EX = 0x64006400; // 0x6400 = FP16 的 1024.0
// Guarantee that the `(a & b) | c` operations are LOP3s.
int lo = lop3 < (0xf0 & 0xcc) | 0xaa > (q, LO, EX); // immLut = 0xea
int hi = lop3 < (0xf0 & 0xcc) | 0xaa > (q, HI, EX);
// We want signed int4 outputs, hence we fuse the `-8` symmetric zero point
// directly into `SUB` and `ADD`.
const int SUB = 0x64086408; // 1024 + 8
const int MUL = 0x2c002c00; // 1/16
const int ADD = 0xd480d480; // -72
typename ScalarType<half>::FragB frag_b;
frag_b[0] = __hsub2(*reinterpret_cast<half2*>(&lo),
*reinterpret_cast<const half2*>(&SUB));
frag_b[1] = __hfma2(*reinterpret_cast<half2*>(&hi),
*reinterpret_cast<const half2*>(&MUL),
*reinterpret_cast<const half2*>(&ADD));
return frag_b;
}
读法是这样的:0x6400 是 FP16 里的 1024.0,它的尾数低位刚好空着。lop3 用一条指令完成 (q & LO) | EX,于是每个 nibble 直接落进尾数,得到 FP16 数值 ;再减掉 SUB = 1024 + 8(0x6408)就是 ,对称零点顺手就减掉了。HI 那一路取的是高位 nibble,数值变成 ,所以要先乘 1/16(0x2c00)再加 -72(0xd480)—— ,一条 __hfma2 搞定。
这套写法的源头不是 Marlin,而是 NVIDIA FasterTransformer 的
interleaved_numeric_conversion.h(vLLM 在vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:161起的注释里直接给出了那两段的 URL 与行号);Marlin 与 vLLM 沿用并微调了它。BF16 版本(vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:192)常数换成EX = 0x43004300(BF16 的128.0)、MUL = 0x3F803F80(1.0)、ADD = 0xC308C308(-136):。它的取 nibble 方式也略有不同——不像 FP16 版那样用LO/HI两个掩码,而是同一个MASK = 0x000f000f配一次q >>= 4,因为 BF16 的尾数只有 7 位,装不下16n。同样的常数在vllm-0.8.5/csrc/quantization/marlin/dense/marlin_cuda_kernel.cu:97、vllm-0.8.5/csrc/quantization/marlin/sparse/common/mma.h:128、vllm-0.8.5/csrc/quantization/marlin/qqq/marlin_qqq_gemm_kernel.cu:142里反复出现,可以对照着看。
比起朴素写法,它省掉的是 16 次 __int2half_rn 那样的"整数→浮点"硬件转换指令,整个解包只剩几条纯位运算 + 两条 half2 算术(同样以 nvcc 13.4、sm_90a 核对:2 条 LOP3.LUT + 1 条 HADD2 + 1 条 HFMA2,没有 I2F)——在 GEMM 内层每秒要跑几亿次的地方,这个差别是决定性的。
类似的技巧在 cutlass / Marlin / Machete 源码里大量存在。这是 LLM 量化 kernel 性能极限的关键之一。
9.7 Marlin INT4 GEMM 简介
Marlin 是 Frantar 等人 2024 年发布的开源 W4A16 GEMM 实现(论文见 arXiv:2408.11743,MARLIN: Mixed-Precision Auto-Regressive Parallel Inference on Large Language Models)。它的目标很明确:在 batch 不为 1 但也不大(约 16–32)的区间里,把"INT4 权重带来的 4 倍带宽节省"尽可能完整地兑现成 4 倍加速(README 在 NVIDIA A10 上的测试:batch 到 16–32 仍接近理想的 4 倍)——朴素的 W4A16 kernel 在 batch 一大就会被解包和调度开销吃掉收益。
Marlin 的核心创新:
- 多级 async pipeline:注意它是 Ampere 级的指令栈——
cp.async.cg.shared.global(源码里包成cp_async4,vllm-0.8.5/csrc/quantization/gptq_marlin/marlin.cuh:71;配对的cp.async.commit_group/cp.async.wait_group在 :82 / :87,主循环里调的是cp_async_wait<stages - 2>(),vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:999)+ldmatrix.sync.aligned.m8n8.x4.shared.b16(vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:129)+mma.sync.aligned.m16n8k16(同文件 :105 的 f16 版与 :112 的 bf16 版),不是 TMA + WGMMA。靠stages级流水让拷贝和计算重叠。 - 快速 INT4→FP16 解包:上面 9.6 节的 lop3 + sub.f16x2 技巧。
- Group scale 走"SMEM 暂存 → 寄存器 fragment":scale 和 A、B tile 一起按流水级用
cp_async4拉进sh_s(vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:841划出sh_s,非 act-order 的分组路径在同文件 :942 发拷贝),再由fetch_scales_to_registers(同文件 :1040)读进双缓冲的frag_s[k % 2]——内层循环算第 k 步时预取第 k+1 步的 scale,SMEM 读取和 MMA 重叠。 - Tile 配置精调:按 SMEM / 寄存器预算精确选择 M/N/K tile 与 stage 数。
真正用上 Hopper TMA + WGMMA 的是它的后继者 Machete(vllm-0.8.5/csrc/quantization/machete/machete_mainloop.cuh:3 的文件头注释就写明它是从 CUTLASS 的 sm90_mma_tma_gmma_rs_warpspecialized_mixed_input.hpp 改出来的)——想看 Hopper 版的混合精度 GEMM 应该读它。
vLLM 在 csrc/quantization/gptq_marlin/ 下集成了它,是 GPTQ / AWQ INT4 权重的默认执行路径之一。收益的上界就是 9.1 节那张表:权重从 FP16 变 INT4,每 token 的 HBM 流量降到 1/4,带宽 bound 的解码延迟也随之接近 1/4。
9.8 FP8 GEMM 的不同路径
Hopper 的 FP8 是另一条路。FP8 不需要 dequant——Tensor Core 直接吃 FP8 输入,但在 Hopper 上必须走 WGMMA:
// Hopper (sm_90a): FP8 走 wgmma, SASS 为 QGMMA.64x8x32.F32.E4M3.E4M3
asm("wgmma.mma_async.sync.aligned.m64n8k32.f32.e4m3.e4m3 ...");
// 同一条 FP8 mma.sync 在 sm_89 上是原生 QMMA.16832.F32.E4M3.E4M3,
// 在 sm_90a 上却被拆成 F2FP.F16.E4M3.UNPACK_B + HMMA.16816 (FP16 速率)
asm("mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 ...");
(以上 SASS 为 nvcc 13.4 编译所得,不同 CUDA 版本可能不同;mma.sync 的 FP8 形式要求 sm_89 及以上,sm_80 上 ptxas 直接报错。)vLLM 在 Hopper 上的 FP8 GEMM 也正是 CUTLASS 3.x 的 TMA + WGMMA 路径(vllm-0.8.5/csrc/quantization/cutlass_w8a8/c3x/scaled_mm_sm90_fp8_dispatch.cuh:21 的 KernelTmaWarpSpecializedPingpongFP8FastAccum)。
FP8 GEMM 的关键:
- 没有 dequant 开销:直接吃 FP8。
- 算力翻倍:H100 SXM5 上 FP8 稠密峰值 1979 TFLOPs vs FP16/BF16 稠密 989 TFLOPs(只有 WGMMA 路径拿得到)。
- 至少每张量一个 scale:FP8 动态范围窄(E4M3 最大 448,E5M2 最大 57344),必须配 scale 防溢出。
NVIDIA 的 Transformer Engine 库提供了"自动 FP8"——它会监控每层激活的范围,动态调整 scale。Megatron-LM 用 TE 训练 FP8 模型。
INT4 vs FP8 的选择:
- 极致带宽节省、低 batch 推理:选 INT4 (W4A16),4 倍带宽节省。
- 训练或大 batch 推理:选 FP8 (W8A8),无 dequant 开销,算力翻倍。
vLLM 这两种都支持,根据用户指定的量化方案选不同的 kernel 路径。
9.9 量化 Kernel 的工程要点
最后总结量化 kernel 的几个工程要点:
- 永远 fuse dequant 到下游计算:单独 dequant 到 HBM 是浪费。
- 位运算技巧极重要:INT4/INT8 解包要用 lop3、sub.f16x2 等专用指令。
- Group scale 提前进寄存器:scale 随流水级进 SMEM,再在内层循环里双缓冲预取到寄存器,别让 scale 读取挡在 MMA 前面。
- Per-token 量化的 reduce 和 GEMM 融合:W8A8 时激活也要量化,把"per-row max + 量化"和上游算子融合,例如 vLLM 的
rms_norm_dynamic_per_token_quant(vllm-0.8.5/csrc/quantization/fused_kernels/fused_layernorm_dynamic_per_token_quant.cu)把 RMSNorm 与 per-token 量化合成一个 kernel。 - 数值稳定性测试:量化后模型质量要测,不要相信"应该没问题"——LLM 量化偶尔会让某些 prompt 输出退化。
9.10 这一章的小结与下一篇
第二篇我们走完了 LLM 推理中的"小算子"舞台:
- 第 5 章 Reduction:所有归约的祖师爷。
- 第 6 章 Online Softmax:让 softmax 流式化的数学魔法。
- 第 7 章 LayerNorm/RMSNorm:用 Welford 把方差变成 1-pass。
- 第 8 章 Element-wise 融合:把无数小算子拼成大 kernel。
- 第 9 章 量化 Kernel:用 INT4/FP8 解决带宽瓶颈。
读完这五章,读者已经具备了 LLM 推理中绝大多数"非 GEMM 非 attention"算子的优化能力。剩下的就是两个真正的大头:GEMM 和 Attention。
第三篇(第 10-13 章)我们正式进入 GEMM——从朴素 GEMM 出发,经过 Tiled GEMM、Tensor Core GEMM,最后到 CUTLASS 设计哲学。读完第三篇,读者会理解为什么 cuBLAS 的 SGEMM 能比教科书式的三重循环快一个数量级以上,以及现代 GEMM kernel 的所有"模板武器"是怎么组装的。
本章动手练习:
- 实现一个 INT8 dequant kernel,对比朴素版本 vs 用 vectorized load + lop3 的优化版本,看带宽差距。
- 用 vLLM 跑同一个模型的 FP16 和 INT4 (Marlin) 版本,记录单 token 延迟和总吞吐。
- 阅读 vLLM 的
vllm-0.8.5/csrc/quantization/gptq_marlin/gptq_marlin.cu:169,对照 9.6 节的推导逐个常数验一遍;再看 :192 的 BF16 版本,想清楚它为什么必须换一套常数。