权重量化通常分为两个阶段:离线阶段把 FP16 权重压缩成 INT4,并保存 scale、zero-point 等量化元数据;本节关注的是另一阶段:推理时如何高效地使用这些 INT4 权重。
所谓 W4A16,指的是权重(Weight)以 4-bit 整数存储,而激活值(Activation)仍保持 16-bit 浮点(FP16/BF16)的推理方案。这是 GPTQ、AWQ 等主流权重量化方法在部署时的标准形态。它的核心矛盾在于:现代 GPU 的 Tensor Core 并不存在 "INT4 × FP16" 这样的混合精度乘法指令。INT4 权重在参与矩阵乘之前,必须先被"反量化"(dequantize)回 FP16。问题是这一步反量化应该发生在哪里?做得不好,量化带来的全部带宽收益会被瞬间抹平;做得好,则能让解码阶段的矩阵乘获得接近 4 倍的加速。
本节先从数据流的角度分析"片上反量化 + GEMM 融合"这一关键设计(2.1.1),再用 Roofline 模型解释 W4A16 为什么在小 batch 下收益巨大、在大 batch 下反而可能变慢(2.1.2),最后动手实现一个完整的 W4A16 反量化矩阵乘 Triton kernel,并与 PyTorch 参考实现对拍验证(2.1.3)。
拿到一份 INT4 权重后,最直观的推理方式是分两步走:
方案 A(朴素两段式): Step 1: dequant_kernel 读取 INT4 权重 + scale/zero,写出完整的 FP16 权重矩阵到显存 Step 2: cublas_gemm 用 cuBLAS 对 FP16 激活 × FP16 权重 做标准矩阵乘
这个方案功能上完全正确,但从访存角度看是一场灾难。以一个 (K,N)=(4096, 4096) 的权重矩阵为例:
Step 1(反量化 kernel):从 HBM 读入约 8 MB 的 INT4 权重 ((4096×4096×0.5)/(1024×1024)=8MB),在 L2/寄存器中完成解包并乘上 scale,再把结果写出 32 MB的FP16权重到 HBM。注意这里必须另开一块显存,而不是原地覆盖。
Step 2(GEMM):cuBLAS 之类的 kernel 再把这 32 MB 的 FP16 权重重新从 HBM 读回来。
从访存流量上讲,一来一回,权重相关的 HBM 流量是 8MB(读 INT4)+ 32MB(写 FP16)+ 32MB(读 FP16)= 72 MB,反而比直接用 FP16 权重(32 MB)多了一倍以上。显存占用上,INT4 权重通常要在整个推理期间常驻(下一次 forward 还要用),因此 8 MB 的 INT4 与 32 MB 的 FP16 临时 buffer 同时存在,峰值 40 MB,比纯 FP16 方案的 32 MB 还要多。
量化省下的存储空间和带宽,在推理路径上被完全浪费了。INT4 的紧凑表示只在 HBM→SM 的那 8 MB 一跳上发挥了作用,紧接着就被还原成 FP16 送回 HBM,后面的 GEMM 面对的完全是一个 FP16 问题。最终 INT4 只剩下"模型文件更小、加载更快"这一点好处。
正确的做法是把反量化融合进 GEMM kernel 内部,让 FP16 形态的权重只存在于GPU的片中,从始至终不落回HBM:
方案 B(融合式,W4A16 的标准做法):
HBM ───── INT4 packed 权重, 0.5bytes/元素 ──▶ Register/Shared Memory
│
片上反量化
w = (q - zero) * scale
│
▼
HBM ───── FP16 激活值 ─────────────────────▶ tl.dot / Tensor Core (FP16 × FP16)
│
▼
FP16 结果写回 HBM
在融合方案中,每个线程块在 K 维上分块循环时,做的事情是:
从 HBM 加载一小块打包的 INT4 权重(通常 8 个 4-bit 值打包进一个 int32);
在寄存器中用移位和按位与指令把 4-bit 数解包出来;
加载对应分组(group)的 scale 和 zero,在片上完成 w = (q - zero) × scale;
把反量化得到的 FP16 小块直接喂给 tl.dot(底层映射到 Tensor Core 的 MMA 指令);
累加进 FP32 累加器,进入下一个 K 分块。
这样一来,权重在 HBM 侧的流量就只有 INT4 本身(外加占比很小的 scale/zero 元数据),相比 FP16 GEMM 减少到约 1/4。反量化的计算成本(移位、减法、乘法)被"藏"进了原本就在等待访存的空闲计算周期里——这正是访存受限 kernel 融合的经典收益模式。
(1)位打包布局(bit-packing layout)。 4-bit 不是任何硬件的原生存储粒度,必须打包进 int8 或 int32。常见做法是沿 K 维把 8 个连续的 4-bit 量化值塞进一个 int32:第 k 个值占据第 (k mod 8) 个 nibble(4-bit 槽位)。这样打包后的权重张量形状从 (K, N) 变为 (K/8, N)。kernel 内解包只需两条指令:
q[k, n] = (packed[k // 8, n] >> ((k % 8) * 4)) & 0xF
生产级实现(如 Marlin、ExLlamaV2)会进一步设计更复杂的交织布局(interleaved layout),让解包出来的数据恰好符合 Tensor Core MMA 指令要求的寄存器排布,省掉片上的重排开销。本章的教学实现采用上面这种最直观的顺序布局。
(2)分组量化元数据的读取。 为了控制量化误差,W4 通常按 group_size = 128 沿 K 维分组,每组独立拥有一对 (scale, zero)。融合 kernel 在处理第 k 行权重时需要索引到第 k // group_size 组的元数据。元数据本身很小(每 128 个权重共享 2 个标量),但要注意让 BLOCK_K 与 group_size 对齐(例如 BLOCK_K = 64、group_size = 128),避免一个 K 分块横跨两个 group 导致的分支处理。
(3)反量化指令与 MMA 的流水重叠。 解包用的是 INT32 ALU,矩阵乘用的是 Tensor Core,二者是 SM 上不同的执行单元。写得好的 kernel 会让"下一块权重的加载与解包"和"当前块的 MMA 计算"重叠起来。在 Triton 中,编译器的自动流水通常能帮我们完成大部分工作。
W4A16 的收益并非无条件成立。利用Roofline 模型可以精确刻画它的适用边界,结论是:W4A16 是为访存受限(memory-bound)场景设计的技术;随着 batch 增大、GEMM 逐渐转入计算受限(compute-bound)区间,它的收益会衰减,甚至由于额外的反量化开销而变为负收益。
LLM 自回归解码时,每一步每个请求只处理 1 个 token。对于单请求(M = 1)的一层线性层 (1, K) × (K, N),其算术强度(Arithmetic Intensity, AI)为:
其中 bw 是权重每元素的字节数。当 K、N 都在几千的量级时,激活和输出的流量可以忽略,访存几乎完全由权重主导:
FP16 权重(bw = 2):AI ≈ 1 FLOP/Byte;
INT4 权重(bw = 0.5):AI ≈ 4 FLOPs/Byte。
而以 A100 为例,其"屋脊拐点"位于 312 TFLOPS ÷ 2.0 TB/s ≈ 156 FLOPs/Byte。无论 FP16 还是 INT4,解码 GEMV 的算术强度都远远低于拐点,深陷屋檐(带宽斜线)之下。在这个区间里,kernel 的执行时间 ≈ 访存字节数 ÷ 带宽,与计算量几乎无关。因此把权重字节数压到 1/4,理论上就能把这条 GEMV 的延迟压到约 1/4——这正是 W4A16 在低并发在线推理中大受欢迎的根本原因。反量化增加的那点整数指令,在带宽瓶颈面前完全"免费"。
当 batch(或 prefill 的序列长度)为 M 时,(M, K) × (K, N) 的算术强度变为:
权重只需从 HBM 读取一次,却被 M 行激活复用,所以算术强度近似随 M 线性增长。令 AI(M) 达到硬件拐点 156,可以解出进入计算受限区的临界 batch:
| 权重精度 | 近似算术强度 | 到达 A100 拐点的临界 M |
|---|---|---|
| FP16(bw = 2) | ≈ M | M ≈ 156 |
| INT4(bw = 0.5) | ≈ 4M | M ≈ 39 |
这张表揭示了一个关键事实:INT4 让 kernel 以 4 倍的速度逼近计算受限区。当 M 超过临界值后,时间瓶颈从 HBM 带宽切换为 Tensor Core 吞吐——而 W4A16 的 MMA 仍然是 FP16 × FP16,峰值算力与普通 FP16 GEMM 完全相同,省带宽不再省时间。更糟的是,此时片上反量化的移位/乘加指令、以及为解包引入的非理想数据布局,开始与主计算争抢发射槽和寄存器,使得 W4A16 kernel 在大 M 下往往慢于 cuBLAS 的纯 FP16 GEMM。
这一分析直接解释了主流推理框架的几个设计决策:
decode 与 prefill 分而治之。 decode(小 M)走 W4A16 融合 kernel;prefill 或大 batch 场景(大 M)有些框架会退回"先整块反量化、再调 cuBLAS"的两段式,或使用 Marlin 这类专门把大 M 区间也优化到接近 FP16 GEMM 水平的 kernel。
量化的主要卖点是延迟和显存,而非吞吐。 在高并发、大 batch 的吞吐型服务里,GEMM 本就接近计算受限,W4A16 的加速空间有限;它真正的战场是显存吃紧的单卡部署和对首 token/逐 token 延迟敏感的在线场景。
优化目标应随 M 漂移。 写 W4A16 kernel 时,小 M 区间应最大化访存效率,大 M 区间则要把反量化开销从关键路径上藏起来。8.6.3 的教学 kernel 以正确性和可读性优先,性能调优的方向会在末尾指出。
下面实现一个完整可运行的 W4A16 流水线,包含四个部分:① 权重量化与位打包、② PyTorch 参考反量化、③ 融合反量化的 Triton GEMM kernel、④ 对拍测试与简单性能测量。
约定与 2.1.1 一致:8 个 4-bit 值沿 K 维打包进一个 int32;按 group_size = 128 沿 K 维分组做非对称量化,量化值 q ∈ [0, 15],反量化公式为 w = (q - zero) * scale。
import torch import triton import triton.language as tl def quantize_pack_w4(w: torch.Tensor, group_size: int = 128): """把 FP16 权重 w (K, N) 量化为 INT4 并沿 K 维打包进 int32。 返回: packed: (K // 8, N) int32,每个 int32 装 8 个 4-bit 量化值 scales: (K // group_size, N) float16 zeros: (K // group_size, N) int32 """ K, N = w.shape assert K % group_size == 0 and K % 8 == 0 wg = w.float().reshape(K // group_size, group_size, N) w_max = wg.amax(dim=1, keepdim=True) w_min = wg.amin(dim=1, keepdim=True) scales = ((w_max - w_min) / 15.0).clamp(min=1e-8) # 4-bit → 16 个量化格点 zeros = torch.round(-w_min / scales).clamp(0, 15) # 非对称零点 q = torch.round(wg / scales + zeros).clamp(0, 15).to(torch.int32) q = q.reshape(K, N) # 位打包:第 k 个值占据 packed[k // 8] 的第 (k % 8) 个 nibble packed = torch.zeros(K // 8, N, dtype=torch.int32, device=w.device) for i in range(8): packed |= q[i::8, :] << (4 * i) return (packed, scales.reshape(K // group_size, N).to(torch.float16), zeros.reshape(K // group_size, N).to(torch.int32)) def dequantize_ref(packed, scales, zeros, group_size: int = 128): """参考实现:解包 + 反量化,还原出完整 FP16 权重,用于对拍。""" Kp, N = packed.shape K = Kp * 8 q = torch.empty(K, N, dtype=torch.int32, device=packed.device) for i in range(8): q[i::8, :] = (packed >> (4 * i)) & 0xF s = scales.repeat_interleave(group_size, dim=0).float() # (K, N) z = zeros.repeat_interleave(group_size, dim=0) # (K, N) return ((q - z).float() * s).to(torch.float16)
两点说明。其一,打包循环里 q[i::8, :] 取出的是第 i, i+8, i+16, … 行,恰好对应每个 int32 中的第 i 个 nibble,与 kernel 内 (k % 8) * 4 的移位规则严格互逆。其二,packed 是有符号 int32,最高 nibble 的解包会触发算术右移(符号位扩展),但随后 & 0xF 只保留低 4 位,结果依然正确——这个细节在 Triton kernel 里同样成立。
@triton.jit def w4a16_gemm_kernel( a_ptr, bq_ptr, scales_ptr, zeros_ptr, c_ptr, M, N, K, stride_am, stride_ak, # A: (M, K) fp16 stride_bk, stride_bn, # packed: (K//8, N) int32 stride_sg, stride_sn, # scales: (K//G, N) fp16 stride_zg, stride_zn, # zeros: (K//G, N) int32 stride_cm, stride_cn, # C: (M, N) fp16 GROUP_SIZE: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, ): pid_m = tl.program_id(0) pid_n = tl.program_id(1) offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) offs_k = tl.arange(0, BLOCK_K) n_mask = offs_n[None, :] < N acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) for k0 in range(0, K, BLOCK_K): # 假设 K % BLOCK_K == 0 k = k0 + offs_k # (BLOCK_K,) 当前 K 分块的行号 # ── 1. 加载激活 A 的 FP16 分块 ───────────────────────────── a = tl.load(a_ptr + offs_m[:, None] * stride_am + k[None, :] * stride_ak, mask=offs_m[:, None] < M, other=0.0) # ── 2. 加载打包权重并在寄存器中解包出 4-bit 量化值 ───────── packed = tl.load(bq_ptr + (k[:, None] // 8) * stride_bk + offs_n[None, :] * stride_bn, mask=n_mask, other=0) shift = (k[:, None] % 8) * 4 q = (packed >> shift) & 0xF # (BLOCK_K, BLOCK_N) int32 # ── 3. 加载分组量化元数据,片上反量化 ───────────────────── g = k[:, None] // GROUP_SIZE s = tl.load(scales_ptr + g * stride_sg + offs_n[None, :] * stride_sn, mask=n_mask, other=0.0) z = tl.load(zeros_ptr + g * stride_zg + offs_n[None, :] * stride_zn, mask=n_mask, other=0) b = ((q - z).to(tl.float16)) * s # FP16 权重块,只存在于片上 # ── 4. Tensor Core 矩阵乘并累加 ────────────────────────── acc += tl.dot(a, b) c_mask = (offs_m[:, None] < M) & n_mask tl.store(c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn, acc.to(tl.float16), mask=c_mask) def w4a16_matmul(a, packed, scales, zeros, group_size: int = 128): M, K = a.shape N = packed.shape[1] c = torch.empty((M, N), device=a.device, dtype=torch.float16) BLOCK_M, BLOCK_N, BLOCK_K = 16, 64, 64 assert K % BLOCK_K == 0 and group_size % BLOCK_K == 0 grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(N, BLOCK_N)) w4a16_gemm_kernel[grid]( a, packed, scales, zeros, c, M, N, K, a.stride(0), a.stride(1), packed.stride(0), packed.stride(1), scales.stride(0), scales.stride(1), zeros.stride(0), zeros.stride(1), c.stride(0), c.stride(1), GROUP_SIZE=group_size, BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N, BLOCK_K=BLOCK_K, num_warps=4, num_stages=3, ) return c
这个 kernel 就是 8.6.1 中方案 B 的直接翻译:整个循环里没有任何 FP16 权重被写回全局内存,b 这个反量化结果的生命周期仅限于寄存器与 tl.dot 之间。注意 group_size % BLOCK_K == 0 的断言保证了每个 K 分块不会横跨量化分组,g = k // GROUP_SIZE 在块内其实是常量,Triton 编译器可以据此做冗余加载消除。
b 只存在于 GPU 寄存器中,从不落回全局显存——这正是 8.6.1 所描述的"片上反量化"的核心。
对拍的基准不是"量化前的原始权重",而是"参考实现反量化出的 FP16 权重"——我们要验证的是 kernel 与参考实现在数值上一致,量化本身的精度损失属于 8.2 节的话题,不应混入此处的判定。
def test_w4a16(M=8, K=4096, N=4096, group_size=128, seed=0): torch.manual_seed(seed) device = "cuda" w = torch.randn(K, N, device=device, dtype=torch.float16) * 0.02 a = torch.randn(M, K, device=device, dtype=torch.float16) packed, scales, zeros = quantize_pack_w4(w, group_size) # 参考路径:显式反量化 + torch.matmul w_hat = dequantize_ref(packed, scales, zeros, group_size) ref = a @ w_hat # 被测路径:融合 kernel out = w4a16_matmul(a, packed, scales, zeros, group_size) torch.testing.assert_close(out, ref, atol=2e-2, rtol=1e-2) max_err = (out.float() - ref.float()).abs().max().item() print(f"[PASS] M={M:>4} K={K} N={N} max_abs_err={max_err:.4e}") if __name__ == "__main__": for M in (1, 8, 16, 33, 128): # 覆盖 GEMV、非对齐 M、大 batch test_w4a16(M=M)
容差设为 atol=2e-2 而非机器精度,原因有二:参考路径在 FP16 下做累加(torch.matmul 的 FP16 GEMM 内部虽用 FP32 累加,但反量化本身经过了 FP16 舍入),而 kernel 路径的反量化在 (q - z) 转 FP16 后与 s 相乘,两条路径的舍入位置不同;K=4096 的长累加会把这些微小差异放大。若想收紧对拍,可以让参考路径也严格模拟 kernel 的舍入顺序。
最后用一个简单的基准测试,观察加速比随 M 的变化:
def bench(M, K=4096, N=4096, group_size=128): w = torch.randn(K, N, device="cuda", dtype=torch.float16) * 0.02 a = torch.randn(M, K, device="cuda", dtype=torch.float16) packed, scales, zeros = quantize_pack_w4(w, group_size) t_fp16 = triton.testing.do_bench(lambda: a @ w) t_w4 = triton.testing.do_bench( lambda: w4a16_matmul(a, packed, scales, zeros, group_size)) print(f"M={M:>5} fp16={t_fp16:.3f} ms w4a16={t_w4:.3f} ms " f"speedup={t_fp16 / t_w4:.2f}x") for M in (1, 4, 16, 64, 256, 1024): bench(M)
在 A100 上运行,典型的趋势是:M=1~16 时 W4A16 明显快于 FP16(教学实现约 2~3 倍,生产级 kernel 可逼近理论的 4 倍);M 到几十之后加速比迅速衰减;M 达到数百时被 cuBLAS 反超——与 8.6.2 中"临界 M ≈ 39 附近发生瓶颈切换"的 Roofline 预测吻合。
本节的实现为可读性做了取舍,与 Marlin、ExLlamaV2 等生产级 W4A16 kernel 的主要差距在于:
| 维度 | 教学实现 | 生产级 kernel(Marlin 等) |
|---|---|---|
| scale/zero 加载 | 每个 (k, n) 元素独立加载,依赖编译器消除组内冗余 | 每组显式只加载一次并广播 |
| 位打包布局 | 顺序打包,解包后需隐式重排才能进入 MMA 寄存器布局 | 离线交织布局,解包输出零重排直达 Tensor Core |
| 流水线编排 | 依赖 Triton 自动流水(num_stages=3) | 手工编排"全局内存加载 → 解包 → MMA"三级流水 |
| 大 M 区间性能 | 反量化开销暴露,显著慢于 cuBLAS FP16 | 反量化彻底移出关键路径,接近 FP16 GEMM 吞吐 |
理解了本节的融合结构之后,阅读这些开源 kernel 的源码就只剩下布局细节的差异了。
W4A16 的本质是一笔"用片上计算换 HBM 带宽"的交易:
反量化必须融合进 GEMM、只发生在片上,量化才能兑现为速度
这笔交易只在访存受限区间有利可图,Roofline 模型给出了收益随 batch 衰减的精确边界
对拍测试的正确基准是参考反量化路径,它把"kernel 写没写对"与"量化损失了多少精度"这两个问题干净地分离开来
Triton fused 4-bit dequant GEMM vs cuBLAS FP16 · NVIDIA RTX PRO 6000 Blackwell 96GB · group_size = 128
权重完全驻留 L2(128 MB),带宽数字不代表 DRAM 吞吐。
| M | W4A16 (ms) | FP16 (ms) | Speedup | W4A16 BW | %peak | |
|---|---|---|---|---|---|---|
| 1 | 0.052 | 0.141 | 2.74× | 692 GB/s | 39% | |
| 2 | 0.052 | 0.138 | 2.67× | 692 GB/s | 39% | |
| 4 | 0.052 | 0.138 | 2.65× | 684 GB/s | 38% | |
| 8 | 0.053 | 0.138 | 2.59× | 672 GB/s | 38% | |
| 16 | 0.055 | 0.140 | 2.55× | 651 GB/s | 36% | |
| 32 | 0.071 | 0.147 | 2.08× | 505 GB/s | 28% | |
| 64 | 0.098 | 0.140 | 1.43× | 365 GB/s | 20% | |
| 128 | 0.130 | 0.148 | 1.14× | 275 GB/s | 15% | |
| 256 | 0.168 | 0.162 | 0.96× | 212 GB/s | 12% | |
| 512 | 0.302 | 0.231 | 0.77× | 118 GB/s | 7% |
整个问题都在 L2 内,且 grid 远填不满 188 个 SM,绝对时间接近 launch 开销下限。
| M | W4A16 (ms) | FP16 (ms) | Speedup | W4A16 BW | %peak | |
|---|---|---|---|---|---|---|
| 1 | 0.015 | 0.037 | 2.42× | 147 GB/s | 8% | |
| 2 | 0.015 | 0.019 | 1.25× | 146 GB/s | 8% | |
| 4 | 0.015 | 0.019 | 1.27× | 146 GB/s | 8% | |
| 8 | 0.015 | 0.021 | 1.37× | 146 GB/s | 8% | |
| 16 | 0.016 | 0.021 | 1.35× | 142 GB/s | 8% | |
| 32 | 0.017 | 0.020 | 1.13× | 128 GB/s | 7% | |
| 64 | 0.025 | 0.020 | 0.78× | 88 GB/s | 5% | |
| 128 | 0.027 | 0.023 | 0.85× | 84 GB/s | 5% | |
| 256 | 0.030 | 0.036 | 1.20× | 75 GB/s | 4% | |
| 512 | 0.039 | 0.040 | 1.03× | 58 GB/s | 3% |
唯一真正打到 DRAM 的 shape,也是三组里最有参考价值的一组。
| M | W4A16 (ms) | FP16 (ms) | Speedup | W4A16 BW | %peak | |
|---|---|---|---|---|---|---|
| 1 | 0.136 | 0.391 | 2.87× | 946 GB/s | 53% | |
| 2 | 0.136 | 0.385 | 2.83× | 944 GB/s | 53% | |
| 4 | 0.137 | 0.385 | 2.82× | 943 GB/s | 53% | |
| 8 | 0.138 | 0.385 | 2.80× | 935 GB/s | 52% | |
| 16 | 0.140 | 0.385 | 2.75× | 917 GB/s | 51% | |
| 32 | 0.190 | 0.392 | 2.06× | 676 GB/s | 38% | |
| 64 | 0.181 | 0.400 | 2.21× | 709 GB/s | 40% | |
| 128 | 0.288 | 0.393 | 1.36× | 447 GB/s | 25% | |
| 256 | 0.468 | 0.444 | 0.95× | 275 GB/s | 15% | |
| 512 | 0.831 | 0.677 | 0.82× | 155 GB/s | 9% |
| M | 2048×2048 | 8192×8192 | 8192×29568 |
|---|---|---|---|
| 1 | 2.42× | 2.74× | 2.87× |
| 2 | 1.25× | 2.67× | 2.83× |
| 4 | 1.27× | 2.65× | 2.82× |
| 8 | 1.37× | 2.59× | 2.80× |
| 16 | 1.35× | 2.55× | 2.75× |
| 32 | 1.13× | 2.08× | 2.06× |
| 64 | 0.78× | 1.43× | 2.21× |
| 128 | 0.85× | 1.14× | 1.36× |
| 256 | 1.20× | 0.96× | 0.95× |
| 512 | 1.03× | 0.77× | 0.82× |