共计 1781 个字符,预计需要花费 5 分钟才能阅读完成。
背景痛点
BERT 等 Transformer 模型在推理时面临三大瓶颈:

- 计算复杂度 :自注意力机制(self-attention) 的 $QK^T/\sqrt{d}$ 计算随序列长度呈平方级增长
- 内存占用:FP32 精度下,BERT-base 的模型大小超过 400MB,导致缓存命中率低
- 并行度不足:原生 PyTorch 执行时存在大量细粒度算子,难以充分利用 GPU 的 SIMD 特性
技术方案
量化方案选择
- 动态量化(Dynamic Quantization)
- 推理时实时计算缩放因子(scale factor)
- 适用场景:模型包含大量 LSTM/Linear 等线性运算
-
代码示例:
model = torch.quantization.quantize_dynamic(model, {torch.nn.Linear}, dtype=torch.qint8 ) -
训练后量化(Post-Training Quantization)
- 需校准数据集统计权重分布
- 典型配置:权重 INT8+ 激活 FP16(WA8A16)
- 精度损失通常 <1%(GLUE 基准测试)
ONNX Runtime 优化
- 执行提供器选择
- CUDA EP:最大化 GPU 利用率
- TensorRT EP:支持层融合(layer fusion)
-
启用 auto_pad 优化减少 Padding 计算
-
图优化示例
sess_options = ort.SessionOptions() sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL sess_options.add_session_config_entry('session.dynamic_block_size', '16')
TorchScript 算子融合
实现自定义的 MultiHeadAttention 融合:
@torch.jit.script
def fused_attention(q: Tensor, k: Tensor, v: Tensor):
scale = 1 / math.sqrt(q.size(-1))
attn = (q @ k.transpose(-2, -1)) * scale
attn = torch.softmax(attn, dim=-1)
return attn @ v
完整代码实现
PyTorch 转 ONNX
torch.onnx.export(
model,
dummy_input,
"bert_quant.onnx",
opset_version=13,
input_names=["input_ids", "attention_mask"],
dynamic_axes={"input_ids": {0: "batch", 1: "seq_len"},
"attention_mask": {0: "batch", 1: "seq_len"}
},
do_constant_folding=True
)
量化校准实现
calibrator = torch.quantization.MinMaxCalibrator()
with torch.no_grad():
for data in calib_loader:
outputs = model(**data)
calibrator.collect_stats(outputs) # 统计张量极值
quant_model = torch.quantization.convert(model)
性能测试
| 配置 | 延迟(ms) | 吞吐量(req/s) |
|---|---|---|
| FP32 (T4) | 58 | 220 |
| INT8 (T4) | 19 | 650 |
| INT8 (Xeon 6248) | 210 | 45 |
精度对比(GLUE-MRPC):
– FP32: 88.5%
– INT8: 87.9%
避坑指南
- 动态 Shape 处理
- ONNX 导出时必须显式声明 dynamic_axes
-
避免在模型内部使用 shape[-1]等动态索引
-
多线程竞争
# 每个线程使用独立 Session sessions = [ort.InferenceSession(model_path) for _ in range(num_threads)] -
量化陷阱
- LayerNorm 层需保持 FP16 精度
- 注意力 mask 的 -INF 值需特殊处理
延伸思考
- 如何利用稀疏注意力 (Sparse Attention) 减少计算量?
- 可变长度输入能否通过动态批处理 (dynamic batching) 提升吞吐?
- 在边缘设备上如何实现混合精度 (Mixed Precision) 推理?
正文完
