同一个模型,有人在GPU上跑出2毫秒,有人跑出8毫秒,这不一定说明换卡就能解决。推理阶段每次请求通常是一张图或一段文本,批大小很小,算力利用率上不去。真正拖慢速度的,是计算图被拆成大量小算子后带来的kernel启动开销和显存往返。每执行一个小算子,CPU需要向GPU提交一次命令,GPU调度完成后从全局显存读输入、再写回输出。激活张量在算子之间反复搬运,执行单元却只算了一点乘加。要优化推理,得先把这些非计算时间压缩掉。

其中算子融合和量化加速是两条低成本路径。前者减少算子数量和中间张量,后者压缩数据类型并利用INT8算力。它们不改变模型结构,也不需要额外训练数据,适合在上线前做一轮系统化部署优化。
一、推理慢的账要算到访存和启动开销上
GPU做一次3乘3卷积只读取9个权重,但中间结果要完整写回显存。一个小算子如果计算量很小,它的执行时间可能只有几微秒,而kernel启动加上数据等待常常就要几微秒到几十微秒。当模型里有上百个这样的算子,延迟自然被抬高。比如MobileNet系列为了少算,拆出了大量逐通道卷积和点卷积,单算子计算强度不高,很容易变成访存瓶颈。ResNet里的每个残差块包含卷积、批归一化、ReLU和加法,未优化时会生成多个中间张量,显存带宽被白白消耗。
定位这类问题可以用GPU分析工具查看时间线,例如Nsight Systems里能看到kernel之间的空隙很大,显存读写总线长时间忙碌但计算单元空闲。CPU侧也会频繁出现kernel launch函数调用。对比计算量相近但延迟更低的模型,往往不是算力差异,而是图结构更紧凑。优化方向也就明确:把相邻小算子合并,让一次kernel完成更多计算,减少启动与访存。
还有一个常被忽略的点是同步。某些算子会隐式触发CPU与GPU同步,或引起显存分配释放。推理若走框架默认路径,可能包含多次设备同步,这些同步点比算子本身还耗时。图优化和融合可以减少这些边界,让整段计算在一个流里连续跑完。
二、算子融合:把中间结果留在缓存里
融合最典型的例子是卷积、批归一化和ReLU。训练时批归一化能稳定梯度,但推理时它只是一个线性变换。把BN的参数合并到卷积权重里,可以少掉两个中间张量。假设卷积输出为y,BN计算为 γ*(y-μ)/√(σ²+ε)+β,等价于给卷积权重乘以 γ/√(σ²+ε),偏置改为 β-μ*γ/√(σ²+ε)。融合后仍然是标准卷积,不会改变模型数学结果。
import torch
import torch.nn as nn
def fuse_conv_bn(conv, bn):
w = conv.weight
gamma = bn.weight
beta = bn.bias
mean = bn.running_mean
var = bn.running_var
eps = bn.eps
scale = gamma / torch.sqrt(var + eps)
fused_w = w * scale[:, None, None, None]
fused_b = beta - mean * scale
return fused_w, fused_b
class ConvBnRelu(nn.Module):
def __init__(self, in_ch, out_ch):
super().__init__()
self.conv = nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False)
self.bn = nn.BatchNorm2d(out_ch)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
x = self.conv(x)
x = self.bn(x)
x = self.relu(x)
return x
部署时可以先调用fuse_conv_bn得到融合权重,再加载到一个不带BN的卷积模块里,ReLU可以继续作为卷积后的原地激活执行。这样原本三步操作变成一步卷积加一步原地激活,中间张量减少2个,显存峰值和读写量都下降。
除了卷积链路,element-wise连续操作也适合融合。比如两个张量相加后再做激活,可以写进一个自定义CUDA kernel里,先读两个输入、求和、做激活、写回一个输出,避免先写回求和结果再读出来。Transformer里的注意力计算也常把缩放、遮罩、softmax和dropout合并,FlashAttention就是一个典型思路,把中间注意力矩阵留在SRAM中反复迭代,而不是写到HBM再读回。不同框架对融合的支持程度不同,TensorRT、ONNX Runtime会在图优化阶段自动匹配这类模式,PyTorch则需要手工替换模块或依赖编译器后端。
并非所有算子都该融合。如果融合后单个kernel变得过大,会降低GPU占用率,尤其在多流并发时反而不利。融合的前提是算子之间有数据依赖且计算强度不高,合并后能明显减少访存。对于计算密集的大矩阵乘,更多考虑的是利用Tensor Core和合适的tiling,而不是单纯合并。
三、量化加速:从FP32到INT8的映射与校准
FP32推理占用带宽大,但神经网络的权重和激活通常分布集中,不需要32位精度来表示每一个数。量化把浮点数值映射到整数区间,推理时用整数运算替代浮点运算。以对称量化为例,先统计某个张量的最大绝对值max,再计算scale=max/127,把原始值x映射为q=round(x/scale),反量化x≈q*scale。INT8范围是[-128,127],权重和激活都可以用这种方式存成8位整数。
import torch
def symmetric_quantize(x, bits=8):
qmax = 2 ** (bits - 1) - 1
max_val = x.abs().max()
scale = max_val / qmax
if scale == 0:
scale = 1.0
q = torch.clamp(torch.round(x / scale), -qmax, qmax).to(torch.int8)
dq = q.float() * scale
return q, scale, dq
这个简易实现演示了对称量化的核心步骤,实际部署中还要处理卷积时的整数乘加和scale传播。根据统计范围的不同,还可以用非对称量化,显式计算零点zero_point,用来表示偏斜分布。权重通常逐通道统计,激活逐张量统计,这样精度损失更小。
量化主要有两条路线:训练后量化PTQ和量化感知训练QAT。PTQ不需要重新训练,从验证集里抽几百到几千张样本做校准,统计每一层激活的分布,选择合理的scale。校准很关键,如果激活出现少量离群值,用最大绝对值会放大scale,导致大量小数值被量化到同一个整数,精度下降明显。常用的KL散度校准或百分位截断就是先找到更紧的阈值再映射。QAT则在训练过程中模拟量化误差,让模型适应整数表示,精度通常更高,但成本也更大。大多数上线场景优先尝试PTQ,精度不足再考虑QAT。
量化带来的加速来自两方面。一是模型体积和显存带宽需求降为原来的四分之一左右,在访存敏感的小模型或大批量场景中提升明显。二是现代GPU和NPU提供INT8的专用乘加单元,峰值算力通常是FP32的两到四倍。比如同为Tensor Core,INT8吞吐远高于FP32。但要注意量化后的算子也要经过融合和优化,否则activation量化与反量化频繁切换,反而会引入新的访存和转换开销。
四、端到端调优:让融合和量化叠加生效
实际部署时,融合和量化不是二选一,而应串行叠加。先做图优化和算子融合,把Conv+BN+ReLU、矩阵乘+偏置+激活等模式合并,再把融合后的节点替换为INT8实现。这样能减少量化节点之间的转换,中间张量直接以INT8形式传递,只在网络入口和出口做一次量化与反量化。TensorRT构建INT8引擎时会同时进行图融合和校准,下面给出一个简化流程。
import tensorrt as trt
def build_int8_engine(onnx_path, calib_data):
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)
with open(onnx_path, 'rb') as f:
parser.parse(f.read())
config = builder.create_builder_config()
config.set_flag(trt.BuilderFlag.INT8)
config.int8_calibrator = calib_data
engine = builder.build_engine(network, config)
return engine
这段代码里模型解析完成后,TensorRT会先识别可融合子图,再依据校准器统计的激活范围插入量化和反量化层。最终生成的引擎里,卷积、归一化和激活已经在同一个INT8 kernel中执行,启动次数和显存读写都大幅降低。
以一个参考模型为例,在相同GPU上,原始FP32推理需要8毫秒左右,做Conv+BN+ReLU融合后降到5毫秒上下,再开启INT8可以降到2.4毫秒,端到端延迟下降约70%。Top1精度下降通常在0.3%到0.8%之间,如果模型本身对量化敏感,可以在敏感层保留FP32或改用逐通道量化。不同硬件表现会有差异,但融合和量化叠加的收益方向基本一致。
最后要验证的不只是延迟数字,还有显存占用、吞吐和精度。延迟优化可能牺牲吞吐,比如大kernel会降低多流并发;INT8虽然提速,但对小批量的收益可能小于较大batch。建议在真实请求分布下做压测,同时设置精度回归脚本,用全量测试集对比量化前后输出。只有延迟、吞吐、显存和精度四项都满足要求,才算把推理速度优化真正落地。