量化 量化(quantization)降低模型权重和激活值的精度,让模型更小、更快、更省电。本文件涵盖数值格式、训练后量化、量化感知训练、仅权重量化(GPTQ、AWQ)、激活量化、混合精度,以及键值缓存(KV-cache)量化 一个 70B 参数的模型用 float16 存储需要 140 GB 显存,超过任何单张 GPU 的容量。量化到 INT4 后只需 35 GB(一张 A100 就够),甚至 20 GB(消费级 RTX 4090 加点 offloading)。量化不是锦上添花的优化,而是让大模型部署在经济上可行的根本手段。 核心权衡:精度越低,显存占用越小、吞吐越高、功耗越低,但会引入量化误差(quantization error),可能损害模型质量。量化的艺术就在于把这种损害降到最低。
量化(quantization)降低模型权重和激活值的精度,让模型更小、更快、更省电。本文件涵盖数值格式、训练后量化、量化感知训练、仅权重量化(GPTQ、AWQ)、激活量化、混合精度,以及键值缓存(KV-cache)量化
一个 70B 参数的模型用 float16 存储需要 140 GB 显存,超过任何单张 GPU 的容量。量化到 INT4 后只需 35 GB(一张 A100 就够),甚至 20 GB(消费级 RTX 4090 加点 offloading)。量化不是锦上添花的优化,而是让大模型部署在经济上可行的根本手段。
核心权衡:精度越低,显存占用越小、吞吐越高、功耗越低,但会引入量化误差(quantization error),可能损害模型质量。量化的艺术就在于把这种损害降到最低。
省显存:INT8 比 FP16 小 2 倍,INT4 比 FP16 小 4 倍。对 LLM 而言,模型权重占用了大部分显存,精度减半,显存需求就减半。
提吞吐:精度越低,每秒能做的运算越多。NVIDIA Tensor Core(第 16 章)FP16 相比 FP32 吞吐翻倍,INT8 相比 FP16 再翻倍,INT4 相比 INT8 还能再翻倍。H100 在 FP8 下可达 989 TFLOPS,而 FP32 只有 67 TFLOPS——足足差了 15 倍。
省带宽:LLM 推理通常是**访存带宽受限(memory-bandwidth-bound)**的(第 16 章的 roofline 模型)。瓶颈在于把权重从显存搬到计算单元,而不是算得太慢。权重越小,搬运的字节越少,每秒生成的 token 就越多。这也是为什么量化对 LLM 推理常常能带来近乎线性的加速。
省电:精度越低,每次运算耗能越少。在数据中心规模(成千上万张 GPU)下,这能省下一笔可观的电费。
| 格式 | 位数 | 指数 | 尾数 | 范围 | 用途 |
|---|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | ±3.4×10³⁸ | 训练(黄金标准) |
| TF32 | 19 | 8 | 10 | ±3.4×10³⁸ | Tensor Core 训练(A100+) |
| FP16 | 16 | 5 | 10 | ±65504 | 混合精度训练 |
| BF16 | 16 | 8 | 7 | ±3.4×10³⁸ | 训练(与 FP32 同范围) |
| FP8 E4M3 | 8 | 4 | 3 | ±448 | 前向传播(Hopper+) |
| FP8 E5M2 | 8 | 5 | 2 | ±57344 | 梯度(范围更宽) |
| INT8 | 8 | — | — | -128 到 127 | PTQ 推理 |
| INT4 | 4 | — | — | -8 到 7 | 仅权重量化 |
| INT2/三值 | 2 | — | — | {-1, 0, 1} | 极限压缩 |
FP8 有两个变体:E4M3(4 位指数、3 位尾数,范围窄但精度高)用于前向传播,E5M2(5 位指数、2 位尾数,范围宽但精度低)用于梯度。Transformer Engine(第 16 章)会按每个张量自动切换两者。
BF16 vs FP16:BF16 的指数范围与 FP32 相同(无溢出风险),但尾数精度更低。FP16 精度更高但范围窄(最大 65504),训练时需要 loss scaling。推理用哪个都行;训练时 BF16 更安全。
整数格式没有指数位——它们表示定点值。要在 float 和 int 之间转换,需要一个缩放因子(scale factor),还可以加一个零点(zero point):x_{\text{float}} = \text{scale} \times (x_{\text{int}} - \text{zero\_point})。
scale 决定分辨率:\text{scale} = \frac{x_{\max} - x_{\min}}{q_{\max} - q_{\min}}。对于 INT8:q_{\min} = -128,q_{\max} = 127。
**对称量化(symmetric quantisation)**令 \text{zero\_point} = 0,于是 \text{scale} = \frac{\max(|x|)}{127}。更简单更快(推理时不用减零点)。
**非对称量化(asymmetric quantisation)**用非零的 \text{zero\_point} 处理不对称分布(例如 ReLU 的输出全是非负数)。把 [x_{\min}, x_{\max}] 映射到无符号 INT8 的 [0, 255]。
Min-max:用观察到的最小/最大值定 scale。简单,但对离群值敏感(一个极端值就把大部分量化范围浪费在用不到的值上)。
百分位(percentile):用 99.99 百分位代替绝对最大值。裁掉极端离群值,给多数值更好的分辨率。被裁掉的值会饱和到 q_{\min} 或 q_{\max}。
MSE 最优:找到使原始张量与量化张量之间均方误差最小的 scale。这是一维优化(搜索可能的裁剪值),通常能给出最好的 PTQ 精度。
基于熵(KL 散度):找到使原始分布与量化分布之间 KL 散度最小的 scale。TensorRT 的 INT8 校准用的就是这个。
# 用 PyTorch 做简化版 PTQ(概念演示) import torch def quantise_tensor_symmetric(tensor, bits=8): qmax = 2 ** (bits - 1) - 1 # INT8 时为 127 scale = tensor.abs().max() / qmax quantised = torch.clamp(torch.round(tensor / scale), -qmax, qmax).to(torch.int8) return quantised, scale def dequantise(quantised, scale): return quantised.float() * scale # 量化一个权重矩阵 weight = torch.randn(512, 512) # 预训练权重 weight_q, scale = quantise_tensor_symmetric(weight, bits=8) weight_reconstructed = dequantise(weight_q, scale) # 量化误差 error = (weight - weight_reconstructed).abs().mean() print(f"Mean absolute error: {error:.6f}") print(f"Compression: {weight.numel() * 4 / (weight_q.numel() * 1 + 4):.1f}x") # +4 字节存 scale
模型在训练过程中学会对量化噪声"免疫"。QAT 通常能恢复 PTQ 损失的大部分甚至全部精度,在低比特宽度(INT4、INT2)下尤其明显。
代价:QAT 需要重训(或微调)模型,对大模型来说很贵。一个 70B 模型做 QAT 可能要花 1 万到 10 万美元算力。而 PTQ 几乎不花钱(只需校准)。
什么时候用 QAT:PTQ 质量无法接受时(通常是 INT4 或更低)、要部署到延迟预算很紧的端侧设备时,或者模型会被量化部署千百万次时(一次性的 QAT 成本会被摊薄)。
关键洞察:量化第 j 列会引入误差。GPTQ 立刻调整剩余所有列来补偿,使整层输出(XW)的变化尽可能小。这就是把**最优脑量化(optimal brain quantisation,OBQ)**用到 transformer 上。
GPTQ 配 4 比特组量化(组大小 128),在大多数 LLM 上困惑度退化 <1%。单个 GPU 上量化一个 70B 模型大约需要 1 小时。
AWQ(激活感知权重量化,Activation-Aware Weight Quantisation,Lin 等,2023)观察到:有少量权重通道(1-3%)远比其他通道重要——它们对应幅值大的激活通道。保护这些"显著通道"能极大降低量化误差。
AWQ 在量化前把这些重要通道放大 s 倍(让它们更大,从而更不受取整影响),同时把对应激活缩小 1/s(保持输出不变)。scale s 按组优化,以最小化整体量化误差。
AWQ 比 GPTQ 更简单(不用算 Hessian)、运行更快、质量相当。它已经成为很多开源 LLM 量化流程的默认选择。
GGUF(GGML Universal Format)是 llama.cpp 用于 CPU 推理的格式,支持很多量化方案:
带 "K" 的变体(k-quants)给重要权重块分配更多比特,思路和 AWQ 类似,但在格式层面实现。Q4_K_M 是大多数模型的甜点:平均 4 比特,质量损失极小。
QuIP(Chee 等,2023)引入了不相干性处理(incoherence processing):量化前用一个随机正交变换旋转权重矩阵。这把信息摊到所有权重上,防止少数离群权重主导量化误差。
直觉:如果有一个权重是 100,其余都约等于 1,用同一个 scale 量化就会把大部分 INT4 范围浪费在那个离群值上。经过一次正交旋转(保持矩阵数学性质不变)后,所有权重幅值相近,均匀量化就好用多了。
QuIP# 在此基础上加了格码本(lattice codebooks):不再映射到均匀整数网格,而是映射到最优格(8 维的 E8 格)上的点。格码能在相同比特数下塞进更多量化点,达到比均匀量化更好的率失真性能。QuIP# 在 2 比特精度下仍能可用——比典型 INT4 方法少用一半比特。
SpQR(Dettmers 等,2023)观察到:极少数权重(0.1-1%)是离群值,对输出质量的贡献不成比例地大。与其把所有权重量化到同一精度,SpQR:
结果:约 99% 的权重被激进量化(小),而关键的 1% 保留全精度(准)。稀疏的离群值存储开销很小(<总大小的 5%)。
HQQ(Half-Quadratic Quantisation,半二次量化,Badri & Shaji,2023)是一种**零样本(zero-shot)**权重量化方法,完全不需要校准数据。它把量化建模成一个半二次优化问题,迭代地求解最优的量化权重和 scale。
优势:不需要校准集意味着没有数据依赖、即拿即量化、也不用担心校准数据不匹配。HQQ 特别适合那些拿不到有代表性校准数据、或数据敏感的模型。
BitNet(Wang 等,2023)把量化推到极致:权重是三值(\{-1, 0, +1\}),每权重只需约 1.58 比特。矩阵乘法变成了只有加减法——根本不需要浮点乘法。
BitNet b1.58(Ma 等,2024)把每个权重限制为 \{-1, 0, +1\}。"1.58 比特"来自 \log_2(3) \approx 1.58。在这个精度下,70B 模型只需约 15 GB 显存,推理完全不用乘法——只需加减和符号翻转。
矩阵乘法变成:
| 格式 | 共享指数 | 每元素位数 | 合计(每元素) | 等价于 |
|---|---|---|---|---|
| MXFP8 | 每块 8 比特 | 8(E4M3/E5M2) | ~8 | 类 FP8 但范围更好 |
| MXFP6 | 每块 8 比特 | 6 | ~6.5 | 介于 FP8 和 INT4 之间 |
| MXFP4 | 每块 8 比特 | 4 | ~4.5 | 类 INT4 但行为更像浮点 |
| MXINT8 | 每块 8 比特 | 8(整数) | ~8.5 | 带共享缩放的 INT8 |
在 NVIDIA Hopper 和 Blackwell GPU 上,用 FP8 训练(不只是推理)已经可行。配方如下:
前向传播:权重和激活用 E4M3(精度更高、范围更窄)。Transformer Engine 用延迟缩放(delayed scaling)动态算每张量的 scale——跟踪上一轮迭代的统计量,应用到当前迭代。
反向传播:梯度用 E5M2(范围更宽、精度更低)。梯度的取值范围比权重/激活宽,多出的指数位能防止溢出。
主权重:优化器状态用 FP32 维护(和第 6 章标准的 FP16 混合精度训练一样)。FP8 只用于矩阵乘,不用于权重更新。
Loss scaling:FP8 仍然需要,就像 FP16 一样。动态 loss scaler 会调整 scale,把梯度值保持在 FP8 可表示的范围内。
FP8 训练在大多数模型规模上能达到和 BF16 训练相当的质量,吞吐提升约 2 倍。它是 H100 集群上新一代大规模训练的默认选择。
激活值(层与层之间流动的中间张量)也可以量化,从而实现全 INT8 计算(权重和激活都是 INT8,用 INT32 累加)。
动态量化(dynamic quantisation):在运行时根据实际激活值算 scale。更准(能适应每个输入)但有额外开销(每层都要算 min/max 或百分位)。
静态量化(static quantisation):校准时算一次 scale 就固定下来。推理时更快(不用运行时统计),但如果校准数据没代表性就不准。
Per-token 量化:序列里每个 token 单独算一个 scale。这对 LLM 至关重要,因为不同 token 的激活幅值可能天差地别(有的 token 激活比其他大 100 倍)。
激活量化比权重量化更难,因为激活依赖数据(每个输入都变),而权重是固定的。"离群值"问题尤其严重:少数激活通道有极端值(均值 100 倍),用和普通通道相同的 scale 量化就会浪费精度。
SmoothQuant(Xiao 等,2022)用数学手段把量化难度从激活(因离群值难量化)迁移到权重(好量化):激活乘 1/s,权重乘 s,其中 s 用来平衡难度。输出 XW = (X \cdot \text{diag}(s^{-1})) \cdot (\text{diag}(s) \cdot W) 不变。
不是所有层对量化同样敏感。注意力层往往能容忍 INT4,而嵌入层和最终分类层需要更高精度。
敏感性分析(sensitivity analysis):逐层单独量化,测量精度影响。高敏感层给更多比特,不敏感层给更少比特。
Transformer Engine(第 16 章,NVIDIA Hopper)在算子级实现动态混合精度:每个矩阵乘根据张量统计在 FP8 和 FP16 间选择,在保质量的同时最大化吞吐。
一个 70B 模型,80 层、64 头、128 维头,序列长度 128K,FP16 下:2 \times 80 \times 64 \times 128 \times 131072 \times 2 = 330 GB。这已经超过单卡显存了。
KV-cache 量化用 INT8 或 INT4 存缓存的 key/value,代替 FP16。量化误差会沿序列累积(每个新 token 都要 attend 到所有缓存的 K/V),但用 per-channel 或 per-head 量化,退化是可接受的。
KV-cache 量化的好处是乘法级的:它允许更长序列(更多上下文)、更大 batch(更多并发用户)、更快推理(搬缓存占的带宽更少)。这是 LLM 服务中投入产出比最高的优化之一。
import jax.numpy as jnp import jax def quantise_int8(tensor): scale = jnp.max(jnp.abs(tensor)) / 127.0 quantised = jnp.clip(jnp.round(tensor / scale), -127, 127).astype(jnp.int8) return quantised, scale def dequantise(quantised, scale): return quantised.astype(jnp.float32) * scale # 正常分布的权重(训练好的模型典型情况) key = jax.random.PRNGKey(0) weights = jax.random.normal(key, (1024, 1024)) * 0.02 q, s = quantise_int8(weights) recon = dequantise(q, s) print(f"Original: {weights.nbytes / 1024:.0f} KB") print(f"Quantised: {q.nbytes / 1024:.0f} KB ({weights.nbytes / q.nbytes:.0f}x smaller)") print(f"Mean abs err: {jnp.abs(weights - recon).mean():.6f}") print(f"Max abs err: {jnp.abs(weights - recon).max():.6f}") print(f"Relative err: {jnp.abs(weights - recon).mean() / jnp.abs(weights).mean():.4%}")
import jax.numpy as jnp import jax key = jax.random.PRNGKey(42) # 激活:大多数通道正常,2 个通道有 100x 的离群值 activations = jax.random.normal(key, (32, 512)) * 0.1 activations = activations.at[:, 0].set(activations[:, 0] * 100) # 离群通道 activations = activations.at[:, 1].set(activations[:, 1] * 50) # 离群通道 # Per-tensor 量化(整张量一个 scale) scale_tensor = jnp.max(jnp.abs(activations)) / 127.0 q_tensor = jnp.clip(jnp.round(activations / scale_tensor), -127, 127) recon_tensor = q_tensor * scale_tensor # Per-channel 量化(每通道一个 scale) scales_channel = jnp.max(jnp.abs(activations), axis=0) / 127.0 q_channel = jnp.clip(jnp.round(activations / scales_channel), -127, 127) recon_channel = q_channel * scales_channel err_tensor = jnp.abs(activations - recon_tensor).mean() err_channel = jnp.abs(activations - recon_channel).mean() print(f"Per-tensor error: {err_tensor:.6f}") print(f"Per-channel error: {err_channel:.6f}") print(f"Per-channel is {err_tensor / err_channel:.1f}x better") print(f"\nOutlier channels waste {(activations.shape[1] - 2) / activations.shape[1]:.0%} " f"of the quantisation range for {2 / activations.shape[1]:.1%} of channels")
def kv_cache_gb(n_layers, n_heads, d_head, seq_len, bytes_per_elem): return 2 * n_layers * n_heads * d_head * seq_len * bytes_per_elem / 1e9 models = [ ("Llama-7B", 32, 32, 128), ("Llama-70B", 80, 64, 128), ("GPT-4 (est)", 120, 96, 128), ] print(f"{'Model':<15} {'SeqLen':>8} {'FP16 (GB)':>10} {'INT8 (GB)':>10} {'INT4 (GB)':>10}") print("-" * 60) for name, layers, heads, d_head in models: for seq_len in [4096, 32768, 131072]: fp16 = kv_cache_gb(layers, heads, d_head, seq_len, 2) int8 = kv_cache_gb(layers, heads, d_head, seq_len, 1) int4 = kv_cache_gb(layers, heads, d_head, seq_len, 0.5) print(f"{name:<15} {seq_len:>8} {fp16:>9.1f} {int8:>9.1f} {int4:>9.1f}") print()