在深度学习落地过程中,模型体积与推理延迟往往是最大的拦路虎。一个经过充分训练的 ResNet-50 FP32 模型约有 98 MB,而大型语言模型的参数规模更是以 GB 甚至 TB 计。模型量化(Model Quantization)通过降低数值精度,在几乎不损失精度的前提下,将模型压缩数倍并显著加速推理。本文将系统梳理从 FP16 到 INT8 的量化原理、PTQ 与 QAT 两大技术路线,并结合 TensorRT 与 ONNX Runtime 给出可落地的工程实践。
一、为什么要做量化
深度学习模型默认使用 32 位浮点数(FP32)存储权重和激活值。量化技术的核心目标是将这些高精度数值映射到低精度表示,带来以下收益:
模型体积缩减:FP32 每个参数占 4 字节,INT8 仅占 1 字节,体积直接缩减为原来的 1/4。对于边缘设备部署,这决定了模型能否装进有限的存储空间。
推理速度提升:现代 CPU(如 Intel AVX-512 VNNI)和 GPU(如 NVIDIA Tensor Core)对 INT8 运算有专门硬件加速。INT8 推理通常可获得 2-4 倍加速,FP16 在支持 Tensor Core 的 GPU 上可获得 1.5-2 倍加速。
内存带宽降低:读取 1/4 大小的权重意味着内存带宽需求同步下降,这在带宽受限的场景(如批量较小的推理)尤为关键。
功耗优化:整数运算单元比浮点单元能效更高。对于移动端和嵌入式设备,低功耗意味着更长的续航和更低的发热。
当然,量化并非免费午餐。低精度表示的动态范围有限,可能引入精度损失。工程上的核心挑战,正是在速度、体积与精度之间找到最佳平衡点。
二、量化基础原理
2.1 对称量化与非对称量化
量化本质是将连续浮点数映射到离散整数。设原始浮点值为 r,量化后的整数值为 q,反量化的浮点值为 r'。
对称量化假设数据分布关于零对称,仅使用一个缩放因子 S:
q = round(r / S)
r' = S * q
此时零点 Z = 0,实现简单,适合近似对称分布的权重参数。
非对称量化(仿射量化)则同时引入缩放因子 S 和零点 Z:
q = round(r / S + Z)
r' = S * (q - Z)
当数据分布明显偏置(如 ReLU 激活后的值均为非负)时,非对称量化能更充分利用有限的 bit 位宽。
2.2 量化粒度
- Per-Tensor(逐张量):整个张量共享同一组
S和Z,存储开销最小,但各通道分布差异大时精度损失明显。 - Per-Channel(逐通道):卷积的每个输出通道使用独立的缩放因子,能更好适应通道间的分布差异,是卷积网络 PTQ 的常用选择。
- Per-Token(逐 Token):在大语言模型中,每个 Token 的激活值使用独立的缩放因子,可有效处理 Transformer 中不同位置动态范围变化剧烈的问题。
2.3 动态范围确定
量化前需要确定张量的最小值和最大值,从而计算 S。常用方法包括:
- Min-Max:直接取张量绝对值的最大值,简单但易受离群值(outlier)影响。
- Entropy(熵校准):通过最小化量化前后分布的信息损失来选择截断阈值,对异常值更鲁棒,TensorRT 的熵校准器即基于此原理。
三、PTQ:训练后量化
PTQ(Post-Training Quantization)指在模型训练完成后,直接对权重和激活进行量化,无需重新训练。它是工程上最便捷的量化方式。
3.1 静态 PTQ
静态 PTQ 需要准备一小部分代表性数据(几百到几千条样本)作为校准集,在推理前离线统计激活值的动态范围。校准集应覆盖实际业务场景中的数据分布,否则量化后的模型在真实数据上可能表现很差。
TensorRT 提供的熵校准器(Entropy Calibrator)会迭代校准数据,收集各激活层的直方图,然后寻找使 KL 散度最小的阈值作为截断点。这比简单 Min-Max 更能保留分布细节。
ONNX Runtime 的静态量化流程类似:先在校准集上运行原模型记录激活统计信息,然后基于这些统计值计算量化参数,最后生成量化模型。
3.2 动态 PTQ
动态 PTQ 不预先校准,而是在推理时实时计算每个输入批次(batch)的激活范围。优点是无需校准数据,对激活分布变化的适应性更强;缺点是运行时统计范围带来额外开销,速度提升不如静态 PTQ 显著。
3.3 PTQ 的适用场景
对于 CNN、ResNet、MobileNet 等结构,静态 PTQ 通常能在精度损失小于 1% 的情况下获得显著加速。当模型存在大量敏感层(如 LayerNorm、Softmax、Attention)或权重分布极不均匀时,PTQ 可能力不从心,此时应考虑 QAT。
四、QAT:量化感知训练
QAT(Quantization-Aware Training)在训练过程中模拟量化效果,使网络参数适应低精度表示,通常能获得比 PTQ 更高的精度。
4.1 伪量化节点
QAT 的核心是在前向传播中插入伪量化(Fake Quantization) 节点:先将权重/激活量化到目标 bit 位,再反量化回浮点。这样网络在前向过程中"感受"到量化的数值误差,并在反向传播中调整参数以补偿这种误差。
4.2 反向传播的直通估计器
伪量化中的 round 函数梯度几乎处处为零,无法直接反向传播。QAT 使用直通估计器(Straight-Through Estimator, STE),在反向传播时直接将量化层的梯度等同于恒等函数的梯度,即跳过量化操作传递梯度。
4.3 QAT 训练流程
典型的 QAT 流程分为三步:
- 训练全精度基线模型:先得到一个收敛良好的 FP32 模型。
- 插入伪量化节点并微调:在基线模型上插入 fake-quant 节点,使用较小的学习率(通常为原始学习率的 1/100)继续训练少量轮次(如原 epoch 的 10%)。
- 转换并导出量化模型:训练完成后将 fake-quant 节点转换为真正的整数运算节点,导出为 ONNX 或 TensorRT 格式。
QAT 适用于对精度要求极为苛刻的场景,或通过 PTQ 无法达到精度要求的模型。其代价是明显的训练成本。
五、FP16 半精度推理
FP16(IEEE 754 half-precision)使用 16 位表示浮点数:1 位符号、5 位指数、10 位尾数。相比 FP32,存储和带宽减半,且 NVIDIA Volta 及之后的 GPU 配备了 Tensor Core,可在 FP16 输入下执行混合精度矩阵乘法并获得大幅加速。
5.1 自动混合精度训练(AMP)
PyTorch 和 TensorFlow 均支持自动混合精度。AMP 在训练时自动选择 FP16 或 FP32:前向和反向中的大部分计算使用 FP16,而损失缩放(Loss Scaling)防止梯度下溢,权重更新保持在 FP32 以保证收敛稳定。
5.2 FP16 的问题与 BF16
FP16 的指数位仅有 5 位,动态范围受限(约 5.96e-8 到 65504)。在训练深层网络或 Transformer 时,容易出现梯度下溢(underflow)。此外,某些层的值域超出 FP16 表示范围时,需回退到 FP32。
BF16(Brain Floating Point) 由 Google 提出,使用 16 位表示:1 位符号、8 位指数、7 位尾数。BF16 与 FP32 共享相同的指数范围,因此几乎不会发生溢出,但尾数精度降低。NVIDIA Ampere 架构及后续 GPU、Google TPU 已原生支持 BF16,成为大规模模型训练的重要替代方案。
六、INT8 量化的深入细节
6.1 哪些层应该量化
并非所有层都适合 INT8 量化。经验法则如下:
- 适合量化的层:卷积层(Conv)、全连接层(Linear)、大部分激活函数(ReLU 等)。这些层的数值分布相对稳定,对量化误差不敏感。
- 通常保留 FP32 的层:LayerNorm、Softmax、Attention 中的 Scale 操作。这些层通常数值动态范围大或对精度极其敏感,量化为 INT8 往往带来显著精度下降。
- 残差连接:需特别注意量化点与反量化点的匹配,避免微小误差在残差路径中被放大。
6.2 权重量化与激活量化
权重量化通常较为容易:训练后的权重分布相对平滑、静态,且可通过 per-channel 粒度精细控制。量化后的权重在推理前即可离线转换,零运行时开销。
激活量化则更具挑战:激活值依赖输入数据,分布随输入变化,必须通过校准或动态统计确定范围。激活中的异常值(outliers)是量化精度损失的主要来源,也是当前大模型 INT8 量化的研究热点。
6.3 纯整数推理
理想的 INT8 推理应全程避免 FP32 回退:输入为 INT8,卷积/矩阵乘法以 INT8 运算执行,累加器使用 INT32 防止溢出,最终的反量化仅在输出阶段执行一次。TensorRT 和 ONNX Runtime 均支持这种整数-only 推理路径,可最大化硬件加速收益。
七、工程实践
以下给出 TensorRT、ONNX Runtime 和 PyTorch 的量化代码示例。
7.1 TensorRT PTQ(Python)
import tensorrt as trt
import pycuda.driver as cuda
import pycuda.autoinit
import numpy as np
class Int8Calibrator(trt.IInt8EntropyCalibrator2):
def __init__(self, data_loader, cache_file="calibration.cache"):
super().__init__()
self.data_loader = data_loader
self.cache_file = cache_file
self.batch = np.zeros((batch_size, 3, 224, 224), dtype=np.float32)
self.d_input = cuda.mem_alloc(self.batch.nbytes)
self.iterator = iter(data_loader)
def get_batch_size(self):
return self.batch_size
def get_batch(self, names):
try:
data = next(self.iterator)[0].numpy()
if data.shape[0] != self.batch_size:
return None
cuda.memcpy_htod(self.d_input, data.ravel())
return [int(self.d_input)]
except StopIteration:
return None
def read_calibration_cache(self):
try:
with open(self.cache_file, "rb") as f:
return f.read()
except FileNotFoundError:
return None
def write_calibration_cache(self, cache):
with open(self.cache_file, "wb") as f:
f.write(cache)
def build_int8_engine(onnx_path, calibrator):
logger = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(logger)
network = builder.create_network(
1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
)
parser = trt.OnnxParser(network, logger)
parser.parse_from_file(onnx_path)
config = builder.create_builder_config()
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30)
config.set_flag(trt.BuilderFlag.INT8)
config.int8_calibrator = calibrator
profile = builder.create_optimization_profile()
profile.set_shape("input", (1, 3, 224, 224), (8, 3, 224, 224), (32, 3, 224, 224))
config.add_optimization_profile(profile)
return builder.build_engine(network, config)
7.2 ONNX Runtime 量化
from onnxruntime.quantization import quantize_static, quantize_dynamic, CalibrationDataReader
import numpy as np
class DataReader(CalibrationDataReader):
def __init__(self, dataloader):
self.dataloader = iter(dataloader)
self.enum_data = []
def get_next(self):
if len(self.enum_data) == 0:
batch = next(self.dataloader, None)
if batch is None:
return None
self.enum_data.append({"input": batch[0].numpy()})
return self.enum_data.pop()
# 静态量化(需校准数据)
reader = DataReader(calibration_dataloader)
quantize_static(
model_input="model.onnx",
model_output="model_int8_static.onnx",
calibration_data_reader=reader,
quant_format=QuantFormat.QDQ, # QDQ 格式与 TensorRT 兼容
activation_type=QuantType.QInt8,
weight_type=QuantType.QInt8,
)
# 动态量化(无需校准数据,仅权重量化)
quantize_dynamic(
model_input="model.onnx",
model_output="model_int8_dynamic.onnx",
weight_type=QuantType.QInt8,
)
7.3 PyTorch 量化 API
import torch
from torch.quantization import get_default_qconfig, prepare, convert
# 准备模型(以 resnet18 为例)
model = torch.hub.load('pytorch/vision:v0.10.0', 'resnet18', pretrained=True)
model.eval()
# 设置量化配置:fbgemm 用于 x86,qnnpack 用于 ARM
model.qconfig = get_default_qconfig('fbgemm')
# 插入 Observer 和 FakeQuantize 节点
model_prepared = prepare(model)
# 在校准数据上运行,收集统计信息
with torch.no_grad():
for images, _ in calibration_loader:
model_prepared(images)
# 转换为量化模型
model_quantized = convert(model_prepared)
# 保存
model_quantized_scripted = torch.jit.script(model_quantized)
torch.jit.save(model_quantized_scripted, "resnet18_int8.pt")
7.4 精度评估
量化后的模型必须通过完整的精度评估验证可用性:
def evaluate(model, dataloader):
model.eval()
correct = total = 0
with torch.no_grad():
for images, labels in dataloader:
outputs = model(images)
_, predicted = outputs.max(1)
total += labels.size(0)
correct += predicted.eq(labels).sum().item()
return 100.0 * correct / total
fp32_acc = evaluate(fp32_model, val_loader)
int8_acc = evaluate(int8_model, val_loader)
print(f"FP32 精度: {fp32_acc:.2f}%")
print(f"INT8 精度: {int8_acc:.2f}%")
print(f"精度下降: {fp32_acc - int8_acc:.2f}%")
通常,对于 ResNet、MobileNet 等 CNN 模型,静态 PTQ 的 Top-1 精度下降可控制在 1% 以内;而对于 Transformer 模型,注意力机制的存在使量化更具挑战,精度下降可能在 1-3% 之间,需结合 per-token 量化或 QAT 改善。
八、性能参考与总结
下表总结了不同量化方案的典型表现:
| 方案 | 体积缩减 | 推理加速 | 典型精度损失 | 适用场景 |
|---|---|---|---|---|
| FP16 | 2x | 1.5-2x | < 0.1% | GPU 推理、 AMP 训练 |
| INT8 PTQ | 4x | 2-4x | < 1%(CNN) | 边缘部署、通用推理 |
| INT8 QAT | 4x | 2-4x | < 0.5% | 精度敏感场景 |
| 动态 PTQ | 4x(权重) | 1.5-2x | 较低 | 无校准数据、快速验证 |
量化失败时的排查方向:
- 模型过小:参数量过少的模型对量化误差更敏感,QAT 通常比重训练更困难。
- 异常值过多:激活值中的离群点会严重干扰动态范围估计,可尝试逐通道或逐 Token 粒度,或采用 SmoothQuant 等先进方法。
- 敏感层未排除:检查 LayerNorm、Softmax 是否被强制量化,适当保留 FP32 回退。
- 校准数据不匹配:校准集必须反映真实推理分布,随机数据或分布偏移的数据会导致灾难性精度下降。
模型量化是连接算法研发与工程部署的关键桥梁。对于绝大多数 CNN 类视觉模型,ONNX Runtime 或 TensorRT 的静态 PTQ 即可在几乎无损精度的情况下获得数倍加速;对于 Transformer 和大语言模型,则需要更精细的粒度控制甚至 QAT。掌握 PTQ 与 QAT 的技术细节,结合硬件特性选择合适的精度与工具链,是每一位深度学习工程师走向成熟的必经之路。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。