← Back to homepage

2.1 反量化推理(W4A16)

权重量化通常分为两个阶段:离线阶段把 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)。

2.1.1 片上反量化与 GEMM 的融合

朴素方案的陷阱:先反量化,再做 GEMM

拿到一份 INT4 权重后,最直观的推理方式是分两步走:

方案 A(朴素两段式):
  Step 1: dequant_kernel   读取 INT4 权重 + scale/zero,写出完整的 FP16 权重矩阵到显存
  Step 2: cublas_gemm      用 cuBLAS 对 FP16 激活 × FP16 权重 做标准矩阵乘

这个方案功能上完全正确,但从访存角度看是一场灾难。以一个 (K,N)=(4096, 4096) 的权重矩阵为例:

从访存流量上讲,一来一回,权重相关的 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 只剩下"模型文件更小、加载更快"这一点好处。

朴素两段式方案的致命问题:反量化后的 FP16 权重落回 HBM,然后 GEMM 又要重新读取。权重相关的 HBM 流量(8 + 32 + 32 = 72 MB)是直接用 FP16 权重(32 MB)的两倍多,峰值显存(40 MB)同样更高,量化在带宽和容量两个维度上的收益都被完全浪费。

融合方案:让反量化只发生在片上

正确的做法是把反量化融合进 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 维上分块循环时,做的事情是:

  1. 从 HBM 加载一小块打包的 INT4 权重(通常 8 个 4-bit 值打包进一个 int32);

  2. 在寄存器中用移位和按位与指令把 4-bit 数解包出来;

  3. 加载对应分组(group)的 scalezero,在片上完成 w = (q - zero) × scale

  4. 把反量化得到的 FP16 小块直接喂给 tl.dot(底层映射到 Tensor Core 的 MMA 指令);

  5. 累加进 FP32 累加器,进入下一个 K 分块。

这样一来,权重在 HBM 侧的流量就只有 INT4 本身(外加占比很小的 scale/zero 元数据),相比 FP16 GEMM 减少到约 1/4。反量化的计算成本(移位、减法、乘法)被"藏"进了原本就在等待访存的空闲计算周期里——这正是访存受限 kernel 融合的经典收益模式。

融合 kernel 需要处理的三个工程细节

(1)位打包布局(bit-packing layout)。 4-bit 不是任何硬件的原生存储粒度,必须打包进 int8int32。常见做法是沿 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_Kgroup_size 对齐(例如 BLOCK_K = 64group_size = 128),避免一个 K 分块横跨两个 group 导致的分支处理。

(3)反量化指令与 MMA 的流水重叠。 解包用的是 INT32 ALU,矩阵乘用的是 Tensor Core,二者是 SM 上不同的执行单元。写得好的 kernel 会让"下一块权重的加载与解包"和"当前块的 MMA 计算"重叠起来。在 Triton 中,编译器的自动流水通常能帮我们完成大部分工作。

2.1.2 从 I/O 受限到计算受限的转变

W4A16 的收益并非无条件成立。利用Roofline 模型可以精确刻画它的适用边界,结论是:W4A16 是为访存受限(memory-bound)场景设计的技术;随着 batch 增大、GEMM 逐渐转入计算受限(compute-bound)区间,它的收益会衰减,甚至由于额外的反量化开销而变为负收益。

解码阶段:深陷访存受限区,收益接近 4 倍

LLM 自回归解码时,每一步每个请求只处理 1 个 token。对于单请求(M = 1)的一层线性层 (1, K) × (K, N),其算术强度(Arithmetic Intensity, AI)为:

AI = 2KN (FLOPs) / [2K (激活) + b_w · KN (权重) + 2N (输出)] (Bytes) ≈ 2 / b_w (FLOPs/Byte)

其中 bw 是权重每元素的字节数。当 K、N 都在几千的量级时,激活和输出的流量可以忽略,访存几乎完全由权重主导:

而以 A100 为例,其"屋脊拐点"位于 312 TFLOPS ÷ 2.0 TB/s ≈ 156 FLOPs/Byte。无论 FP16 还是 INT4,解码 GEMV 的算术强度都远远低于拐点,深陷屋檐(带宽斜线)之下。在这个区间里,kernel 的执行时间 ≈ 访存字节数 ÷ 带宽,与计算量几乎无关。因此把权重字节数压到 1/4,理论上就能把这条 GEMV 的延迟压到约 1/4——这正是 W4A16 在低并发在线推理中大受欢迎的根本原因。反量化增加的那点整数指令,在带宽瓶颈面前完全"免费"。

batch 增大:交叉点提前到来

当 batch(或 prefill 的序列长度)为 M 时,(M, K) × (K, N) 的算术强度变为:

AI(M) = 2MNK / [2MK + b_w · KN + 2MN] ≈ 2M / b_w (当 M << K, N)

权重只需从 HBM 读取一次,却被 M 行激活复用,所以算术强度近似随 M 线性增长。令 AI(M) 达到硬件拐点 156,可以解出进入计算受限区的临界 batch:

权重精度近似算术强度到达 A100 拐点的临界 M
FP16(bw = 2)≈ MM ≈ 156
INT4(bw = 0.5)≈ 4MM ≈ 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 在 A100 上的收益拐点大约在 M ≈ 39。低于此值时收益接近理论 4×,高于此值后收益急剧衰减。这一结论在 2.1.3 的 benchmark 实验中会得到验证。

2.1.3 代码实战:反量化矩阵乘 Triton 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

第一步:权重量化与位打包(离线,PyTorch 实现)

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 GEMM 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 编译器可以据此做冗余加载消除。

关键设计:反量化后的 FP16 权重 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 的舍入顺序。

第四步:验证 2.1.2 的预言

最后用一个简单的基准测试,观察加速比随 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 预测吻合。

与生产级 kernel 的差距

本节的实现为可读性做了取舍,与 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 带宽"的交易:

在实际系统中,W4A16 融合 kernel 是权重量化推理的核心加速组件。理解本节的"方案 A vs 方案 B"、Roofline 临界 M、以及对拍验证方法论,是阅读 Marlin、ExLlamaV2、BitBLAS 等生产级源码的必要基础。

W4A16 GEMM Benchmark

Triton fused 4-bit dequant GEMM vs cuBLAS FP16 · NVIDIA RTX PRO 6000 Blackwell 96GB · group_size = 128

test1 K=8192, N=8192 · weights 32 MB

权重完全驻留 L2(128 MB),带宽数字不代表 DRAM 吞吐。

MW4A16 (ms)FP16 (ms) Speedup  W4A16 BW%peak
10.0520.1412.74×692 GB/s39%
20.0520.1382.67×692 GB/s39%
40.0520.1382.65×684 GB/s38%
80.0530.1382.59×672 GB/s38%
160.0550.1402.55×651 GB/s36%
320.0710.1472.08×505 GB/s28%
640.0980.1401.43×365 GB/s20%
1280.1300.1481.14×275 GB/s15%
2560.1680.1620.96×212 GB/s12%
5120.3020.2310.77×118 GB/s7%

test2 K=2048, N=2048 · weights 2 MB

整个问题都在 L2 内,且 grid 远填不满 188 个 SM,绝对时间接近 launch 开销下限。

MW4A16 (ms)FP16 (ms) Speedup  W4A16 BW%peak
10.0150.0372.42×147 GB/s8%
20.0150.0191.25×146 GB/s8%
40.0150.0191.27×146 GB/s8%
80.0150.0211.37×146 GB/s8%
160.0160.0211.35×142 GB/s8%
320.0170.0201.13×128 GB/s7%
640.0250.0200.78×88 GB/s5%
1280.0270.0230.85×84 GB/s5%
2560.0300.0361.20×75 GB/s4%
5120.0390.0401.03×58 GB/s3%

qwen2-72b-ffn K=8192, N=29568 · weights 121 MB

唯一真正打到 DRAM 的 shape,也是三组里最有参考价值的一组。

MW4A16 (ms)FP16 (ms) Speedup  W4A16 BW%peak
10.1360.3912.87×946 GB/s53%
20.1360.3852.83×944 GB/s53%
40.1370.3852.82×943 GB/s53%
80.1380.3852.80×935 GB/s52%
160.1400.3852.75×917 GB/s51%
320.1900.3922.06×676 GB/s38%
640.1810.4002.21×709 GB/s40%
1280.2880.3931.36×447 GB/s25%
2560.4680.4440.95×275 GB/s15%
5120.8310.6770.82×155 GB/s9%

Speedup 汇总

M 2048×2048 8192×8192 8192×29568
12.42×2.74×2.87×
21.25×2.67×2.83×
41.27×2.65×2.82×
81.37×2.59×2.80×
161.35×2.55×2.75×
321.13×2.08×2.06×
640.78×1.43×2.21×
1280.85×1.14×1.36×
2561.20×0.96×0.95×
5121.03×0.77×0.82×

© Xiaoyi | Homepage