ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Transformer模型优化实战:从推理延迟到显存占用的全面压降

Transformer模型优化实战:从推理延迟到显存占用的全面压降 1. 从线上事故说起为什么我决定认真对待模型优化上个月生产环境出了一次不小的状况。某个BERT类模型的服务在晚高峰时段P95延迟从80毫秒一路飙到450毫秒GPU显存吃满后开始排队连锁反应导致调用方超时重试最终把数据库连接池也拖垮了。复盘的时候发现模型本身在一次小版本迭代里增加了一个注意力头参数量只涨了不到5%但推理耗时却翻了一倍多——这让我意识到过去那种“模型能跑就行”的思路走不远了。于是就有了这个名为Model-Optimizer的优化项目。核心目标很简单在不重新训练、不动模型整体架构的前提下把推理延迟降下来、显存占用压下去、吞吐提上去同时尽量保住模型精度。优化对象是一个部署在PyTorch框架下的Transformer类文本模型日常处理分类、抽取和排序任务训练时是FP32精度上线后一直是原生PyTorch推理没有任何加速框架加持。先说结论经过三轮迭代优化最终模型推理延迟降低了73%显存占用下降了58%batch大小为32时的吞吐量提升了接近3倍精度损失控制在0.5%以内。整个过程涉及量化、剪枝、蒸馏、部署侧算子融合等多条技术路线中间踩了不少坑也推翻了好几次方案。这篇文章就把完整的优化思路、操作细节和实测数据整理出来给正在做类似事情的同行一个参考。2. 先别急着动手搞清楚瓶颈到底在哪2.1 用Profiler把耗时分布拉出来看优化之前最重要的一件事不是找优化方案而是定位瓶颈。我见过太多人一上来就搞量化、上TensorRT结果发现瓶颈根本不在算子计算上而在数据加载或者CPU-GPU拷贝环节——那种优化做了等于白做。我用的第一招是PyTorch自带的torch.profiler把线上真实请求的推理过程完整profile了一遍。核心代码大致是这样import torch from torch.profiler import profile, ProfilerActivity with profile(activities[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: with torch.no_grad(): for _ in range(100): output model(input_tensor) print(prof.key_averages().table(sort_bycuda_time_total, row_limit20))最关键的输出是各算子的cuda_time_total占比。我当时的模型结构是6层Transformer每层包含多头自注意力、前馈网络、两个LayerNorm和一个残差连接。Profiler给出的时间分配非常清楚自注意力模块QKV投影、注意力分数计算、输出投影占GPU总时长的41%前馈网络两个线性层占30%LayerNorm和残差Add合计占12%其余是Embedding、激活函数和内存拷贝这里有一个非常反直觉的事实LayerNorm这种看起来没什么计算量的算子在Transformer里反而占了不小的耗时比例。原因在于LayerNorm在PyTorch原生实现中会做多次独立的Reduce操作每次都要读写一遍内存带宽消耗远高于计算消耗。这个细节在后续优化中帮我省了不少事。2.2 区分计算瓶颈和访存瓶颈Profiler只能告诉你时间花在哪但你要判断这个“花在哪”是计算量大导致的还是访存/带宽导致的。判断方法很简单把GPU利用率和内存带宽利用率的比值拉出来看。我用了nvidia-smi dmon -c 10观察实时指标再用Nsight Compute做了二次确认。结果发现模型大部分算子属于访存受限memory-bound而非计算受限compute-bound。最典型的证据是注意力分数计算时SM占有率并不高但显存读写压力很大。这个结论直接决定了优化路线——如果模型是计算受限首选是算子融合和低精度推理如果是访存受限那么减少中间张量的读写才是核心突破口。我的模型两者都有但访存问题更突出所以后续优化的优先级排序是算子融合 量化 剪枝 蒸馏。这个排序和网上很多教程的顺序不一样但完全是从实际问题推导出来的。提示不要盲目套用别人的优化顺序。你必须先搞清楚自己的模型是哪种瓶颈类型再做对应策略。判断依据就是Profiler数据加Nsight的occupancy和memory throughput指标。2.3 别忘了检查输入侧的开销还有一个小坑值得单独拎出来说。我最初profile的时候把数据加载和预处理全放在循环外面了导致CPU耗时几乎为零结果漏掉了一个大问题——线上服务里每次请求进来的token长度不固定最长512、最短十几个而模型推理时按最长的那个padding到512。这带来了两个浪费一是无效token的算力白白消耗二是多出来的显存占用。后面我用torch.nn.utils.rnn.pad_sequence按batch内最大长度动态padding配合attention_mask同一批数据里长短差异大的情况下整体推理时间能再缩8%到10%。这是零成本优化里性价比最高的一类操作强烈建议先做。3. 算子融合自己做从LayerNorm动刀3.1 为什么PyTorch原生的LayerNorm这么慢前面提到了LayerNorm的耗时占比这里展开说说它是怎么被优化掉的。Transformer里的LayerNorm通常配合残差连接一起出现计算顺序是x sublayer(x)然后再做LayerNorm。PyTorch原生会把x sublayer(x)先算出结果写入中间张量再把这个张量喂给LayerNormLayerNorm内部又要分几步先算均值、再算方差、再归一化、再乘gamma加beta。每一步都是一次独立的kernel launch每一步都要读写一遍完整的内存。算一笔账一个batch大小为32、序列长度128、hidden size 768的中间张量大小是32×128×768 3.1M个元素FP32下就是12.6MB。整个LayerNorm过程要把这块数据读好几遍、写好几遍单纯带宽就能吃掉一大截时间。3.2 手动融合残差和LayerNorm合并成一个Kernel优化的思路是把“残差Add”和“LayerNorm”合并成一个自定义的融合算子一次读取就能完成全部计算。PyTorch本身有torch.nn.functional.layer_norm但不能直接和残差融合所以要用CUDA扩展或者torch.utils.cpp_extension写一个简单的融合kernel。由于我们用的是偏学术向的原型验证并没有急着上CUDA C写kernel而是先用PyTorch的torch.jit.script做了第一版融合尝试import torch torch.jit.script def fused_residual_layernorm(x, residual, weight, bias, eps: float 1e-5): x x residual mean x.mean(dim-1, keepdimTrue) var x.var(dim-1, unbiasedFalse, keepdimTrue) x_norm (x - mean) / torch.sqrt(var eps) return x_norm * weight bias脚本化的效果比想象中好单算子耗时下降了约35%因为它省去了中间张量的多次写入。但真正的质变还是靠后续量化到INT8之后才出现这个后面细讲。3.3 再接再厉QKV投影的矩阵合并Transformer自注意力模块里有三组投影矩阵——Query、Key、Value。原始实现是三个独立的nn.Linear每个都要做一次大的矩阵乘法中间产生三个独立的输出张量。优化方式是把它们合并成一个更大的线性层把[768, 768]的三个权重矩阵横向拼接成[768, 2304]一次GEMM算出QKV的所有结果再从结果中按列切分。对底层而言一次大矩阵乘法远比三次小矩阵乘法高效尤其是批量场景下能更好地利用GPU的Tensor Core。实现上最省事的是用nn.Linear(768, 2304)替代原来的三个nn.Linear(768, 768)然后在forward里切分class FusedQKVProjection(torch.nn.Module): def __init__(self, hidden_size, num_heads): super().__init__() self.hidden_size hidden_size self.head_dim hidden_size // num_heads self.num_heads num_heads self.qkv_proj torch.nn.Linear(hidden_size, 3 * hidden_size, biasTrue) def forward(self, x): batch, seq_len, _ x.shape qkv self.qkv_proj(x) # [batch, seq_len, 3*hidden] qkv qkv.reshape(batch, seq_len, 3, self.num_heads, self.head_dim) q, k, v qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2] return q, k, v这一步之后注意力部分的耗时进一步下降和原来的三次线性层相比大约又省了20%。整体效果相对原始模型累积下来已经快接近40%的时间缩减。4. 量化这条路从FP16到INT8的实战记录4.1 精度需求决定了下限硬件能力决定了上限量化是我整条优化路线里的重点也确实是收益最大的一步。在动手之前需要先确认两件事模型能不能接受精度损失以及目标GPU平台是否支持INT8的加速。我的模型是文本分类/抽取类精度的敏感度中等上线前评测要求F1不能低于原始版本的98.5%。硬件平台是NVIDIA T4和A10G都支持INT8推理加速而且T4对INT8 Tensor Core的支持非常成熟配合TensorRT使用效果更佳。所以确定量化目标为INT8同时保留FP16的中间层策略。4.2 先上PTQ效果不理想再考虑QAT业界常见的两条路线是PTQ训练后量化和QAT量化感知训练。PTQ省事但精度损失可能较大QAT效果好但需要重训模型成本高。我选择先跑一轮PTQ做基线如果不满足精度预算再决定要不要上QAT。PyTorch官方的量化工具链从2.0开始已经很完善用torch.ao.quantization可以直接做静态量化import torch from torch.ao.quantization import quantize_fx, QConfigMapping from torch.ao.quantization.observer import MinMaxObserver, PerChannelMinMaxObserver model_to_quantize model.eval() qconfig QConfigMapping() qconfig.set_global(torch.ao.quantization.default_qconfig) # 或者手动指定 # qconfig.set_global( # torch.ao.quantization.QConfig( # activationMinMaxObserver.with_args(dtypetorch.quint8, qschemetorch.per_tensor_affine), # weightPerChannelMinMaxObserver.with_args(dtypetorch.qint8, qschemetorch.per_channel_symmetric), # ) # ) prepared_model quantize_fx.prepare_fx(model_to_quantize, qconfig, example_inputs) # 喂入校准数据跑几个batch统计激活值分布 with torch.no_grad(): for i, (inputs, _) in enumerate(calib_dataloader): prepared_model(inputs) if i 20: break quantized_model quantize_fx.convert_fx(prepared_model)这里有一个我最想强调的细节校准数据的选取决定了PTQ的成败。一开始我贪图方便直接拿了训练集里随机抽的500条文本做校准结果精度暴跌F1直接掉了3个多点完全没法用。后来排查原因发现校准数据要尽量覆盖模型在实际部署中会遇到的数据分布。训练集里的文本以新闻和科技类为主但线上实际请求中有大量口语化、带噪声的短文本两者特征分布差异很大。激活值的量化范围如果由校准集决定那就必须让校准集代表线上真实分布。最后我改用线上采样日志里的2000条真实请求文本重新做校准精度损失立刻回到了可接受范围。4.3 PTQ之后该量化哪些层不该量化哪些层标准做法是默认量化所有nn.Linear和nn.Conv2dEmbedding和LayerNorm保持高精度。但对于Transformer模型我额外做了几处调整Embedding层不量化因为查询向量对精度极其敏感稍微偏差就会导致检索类任务效果崩坏LayerNorm保持FP32或FP16因为归一化操作涉及除法和小数值运算低精度下误差会累积到后续层注意力Softmax之前的分值计算保持高精度softmax本身对数值范围敏感把这些敏感层排除在INT8之外后精度损失进一步收窄从原来的1.2%降到0.4%左右。代价是这些层仍然以FP16运行占整体计算的比重不大所以总体加速收益依然非常可观。注意量化不是“照着官方教程跑一遍就行”。你必须理解每一层在模型中的职责再决定是否纳入量化范围。尤其是Softmax、LayerNorm、Embedding这类对小数精度敏感或者参与归一化运算的层默认量化往往会造成不必要的精度损失。4.4 INT8的坑per-tensor和per-channel的取舍在配置量化策略时激活值我用的是per_tensor_affine权重用的是per_channel_symmetric。这里有一个常见的误解有人觉得per-channel比per-tensor好就一律用per-channel其实两者有明确的使用场景。权重侧用per-channel更合理因为权重矩阵的每个输出通道数值范围差异可能很大统一用同一个scale会浪费大量表示精度。激活侧则相反一个batch里所有通道共享一个scale因为激活值经过激活函数之后分布相对统一用per-channel反而会增加计算复杂度。实际参数配置可以参考下表配置项设置原因量化精度INT8权重和激活Tensor Core原生支持速度最快权重量化粒度per-channel symmetric保留各通道数值范围精度损失更小激活量化粒度per-tensor affine激活分布相对统一实现简单校准算法MinMax 直方图截断MinMax对长尾分布敏感需配合截断量化范围[0, 255]无符号激活配合affine量化简化推理逻辑排除层Embedding、LayerNorm、Softmax对精度敏感保持高精度推理这里校准算法我用的是PyTorch默认的observer对长尾分布的场景可以换成HistogramObserver来做百分位数截断效果会更稳。不过因为后面还要上TensorRT做融合优化量化方案会被框架重新处理所以PTQ阶段我用默认配置够用了。5. 剪枝不是玄学如何安全地砍掉冗余参数5.1 一个关键决策结构化剪枝 vs 非结构化剪枝量化解决的是“算得快”的问题剪枝解决的是“参数少”的问题。两者结合能进一步压低显存占用也能顺带提升推理效率。剪枝分两大类非结构化剪枝是直接把权重矩阵里的某些元素置零保留稀疏结构结构化剪枝是整列、整行地删除神经元或注意力头。非结构化剪枝的精度保持效果好但除非硬件和推理库原生支持稀疏矩阵加速否则压下来的参数无法转化为实际速度提升。结构化剪枝的精度损失稍大但剪完之后得到的是稠密小模型任何框架都能白捡速度收益。我的目标是部署到生产环境所以只考虑结构化剪枝。合理的方向集中在两个维度注意力头剪枝和前馈网络中间维度剪枝。5.2 用什么标准决定剪谁判断哪个注意力头可以被剪掉最直接的方法是计算每个头对最终输出的“重要性分数”。常用的有基于梯度的方法比如在验证集上做反向传播统计每个头对应梯度的L2范数也有基于启发式的重要性估计比如每个头的输出与最终预测标签的相关性。我实际采用的是后者的一个变体——基于验证集Loss的重要性评估。做法是逐个把某个注意力头以及对应输出投影的行列mask掉在验证集上重新跑一遍观察Loss变化。变化最小甚至反向减少的头就是冗余度最高的优先剪掉。这个方法的缺点是计算量偏大因为每个头都要单独跑一次验证集但6层模型加上128个头跑一次验证集也就几分钟可以接受。前馈网络的中间维度剪枝也类似不过我不是用逐维度遍历的方式而是用torch.nn.utils.prune配合L1范数做channel级别的剪枝——把L1范数最小的中间神经元置零再删除对应的行和列。5.3 剪完必须微调否则精度会崩剪枝和量化不一样的地方在于量化完可以不做任何训练直接用但剪枝完如果不微调精度损失几乎一定会超标。因为被剪掉的头和维度在训练时是参与梯度更新的你强行删掉它们相当于丢掉了原有模型的部分拟合能力必须靠微调补回来。我的微调策略是学习率降到原始训练学习率的十分之一冻结前几层Transformer的参数通常低层学到的是通用语法特征更适合保留只更新后几层和输出头训练1到2个epoch就够了。微调之后剪枝带来的F1损失从原来的1.1%收窄到0.3%以内。一个实际注意点微调完后之前量化时用的校准数据要重新采一批。因为剪枝改变了模型内部激活值的分布旧的量化scale可能不再适配如果沿用旧参数直接上INT8精度会再掉一截而且很难排查。5.4 剪枝和量化的搭配顺序我建议先剪后量实操顺序上我强烈建议先做结构化剪枝、再做量化不要反过来。原因是量化会先把权重变成低精度离散值再做剪枝时重要性判断依据的梯度信号已经被扰动结果更不可靠。先剪后量剪枝时拿到的还是完整的FP32权重重要性评估更准确。实际跑下来的数据也验证了这一点先剪后量的最终F1比先量后剪高0.6%左右。虽然差距不算很大但优化项目本身就是在抠细节这种能白赚的精度损失当然要尽量规避。6. 知识蒸馏让大模型当老师6.1 蒸馏的目标用硬标签之外的软信息“教”小模型经过量化和剪枝之后模型的参数量从原来的52M降到约31M精度损失累计大概0.8%。但线上业务方对精度要求很严F1不能低于原始98.5%也就是说我总共只有1.5%的精度操作空间。此时继续靠剪枝和量化挖潜力的余地越来越小于是我把目光转向了知识蒸馏。蒸馏的思路很简单用原始大模型Teacher的输出分布作为软标签去训练一个小模型Student。软标签里包含了模型对各类别置信度的完整分布信息相比独热硬标签能传递更丰富的语义关系。6.2 蒸馏温度的选择一个经验驱动的问题蒸馏用到温度参数T来控制Teacher输出的软化程度。温度越高输出分布越平滑类别间的差异越小温度越低分布越接近独热。T的取值直接影响Student能从Teacher身上学到多少“暗知识”。我做了几组实验对比温度TStudent最终F1备注197.3%几乎等于直接用硬标签训练398.0%有一点软化信号但还不够598.4%最佳区间精度逼近原始模型897.9%过度平滑类别间区分度下降1296.8%信息过于稀释开始反效果最终确定T5。训练时Loss分两部分与硬标签的交叉熵Loss以及与Teacher软标签的KL散度Loss两者加权求和。权重通常取值0.9给软标签、0.1给硬标签我从这个经验值开始调发现对我这个任务场景来说很稳定没有进一步调整的必要。6.3 Student模型怎么选直接剪枝后的模型就是好起点关于Student的选择有两种常见路径一种是从零训练一个更小的模型结构另一种是直接把剪枝压缩后的模型当作Student用Teacher的软标签再微调一轮。我选的后者因为剪枝后的模型已经具备原始模型的“骨架”蒸馏只需要在它的基础上做知识校准比从零训省时省力。效果也确实不错经过蒸馏微调后约31M参数的小模型F1从97.6%提升到98.3%距离原始模型的98.8%只差0.5%已经进入业务可接受范围。加上量化和算子融合带来的速度收益整体性价比非常高。蒸馏这段实操下来我最大的感受是不要觉得蒸馏是只有大厂才用得起的重型武器。如果手里已经有一个训练好的模型把它当Teacher配合剪枝产出的模型当Student训练成本远低于从零训练一个小模型效果也更可控。7. 部署侧的最后一公里TensorRT与算子融合7.1 为什么PyTorch之外还需要一个推理引擎模型优化做到这一步PyTorch原生推理的速度已经有明显提升但如果还想再狠压一头就需要借助专门的推理引擎。我用的是TensorRT它和PyTorch最大的区别在于PyTorch是动态图执行每个算子独立调度有很多kernel launch开销TensorRT会做整图优化、算子融合、内存复用还能把INT8量化后的算子直接映射到Tensor Core。从PyTorch导出到TensorRT主流路径是先把模型转成ONNX再用trtexec或TensorRT Python API加载ONNX做优化。这里最核心的问题是ONNX导出时要把动态维度固定或者显式声明。7.2 ONNX导出遇到的坑动态形状和动态维度我的模型有两个动态维度batch大小和序列长度。ONNX导出时如果不显式声明TensorRT会默认它们是固定值跑起来要么报错要么必须按最大形状分配显存。正确做法是在导出时用dynamic_axes标明import torch dummy_input torch.randn(1, 128, 768, devicecuda) torch.onnx.export( model, dummy_input, model.onnx, input_names[input_ids, attention_mask], output_names[logits], dynamic_axes{ input_ids: {0: batch, 1: seq_len}, attention_mask: {0: batch, 1: seq_len}, }, opset_version17, )但这里有个后续麻烦动态形状意味着TensorRT要为一个模型编译多种shape的组合优化时间暴涨而且如果某个shape组合没被覆盖到实际请求就可能掉到未优化的模式速度大打折扣。我最终的做法是限制序列长度为固定值batch保持动态。因为文本长度经过动态padding之后最长512无论线上怎么请求输入形状都在一个可控范围内。我把模型按seq_len256和seq_len512分别编译了两份engine根据实际负载路由。7.3 TensorRT能为我们额外省下什么TensorRT在ONNX基础上做了大量算子融合。比如把Linear GELU融合成一个算子把QKV投影 切分优化为一个大GEMM加slice把LayerNorm的多个reduce合并成一次扫描。这些融合动作在PyTorch里即使手动做了也只是冰山一角TensorRT能做得更彻底。实测上同样的INT8模型TensorRT推理比PyTorch的原生INT8推理快了约38%。这个差距主要来自kernel launch次数减少和显存访问优化相当于把“优化后的模型”又往前推了一大步。提示TensorRT版本的兼容性问题必须提前摸清。不同版本对ONNX算子支持的差异很大我最初用的TensorRT 8.4导出的engine在另一台机器的8.2环境上根本加载不了。建议在CI里固定TensorRT版本或者直接用容器镜像绑定。8. 一组实测数据与踩坑清单8.1 各阶段优化效果对比把整个优化链条的每个阶段数据汇总如下方便直观对比优化阶段P95延迟(ms)显存占用(MB)F1说明原始FP32 PyTorch342215098.8%线上事故时的基线动态padding优化310189098.8%小改动大收益算子融合LayerNormQKV合并242167098.8%纯PyTorch script实现PTQ INT8量化128104097.4%校准集优化前PTQ INT8量化校准集修正128104098.2%校准集换成线上真实分布结构化剪枝31M参数9986097.6%剪掉约45%注意力头剪枝后蒸馏微调9383098.3%知识蒸馏补回约0.7%精度TensorRT INT8推理5889098.3%延迟进一步压降约38%最终线上能拿得出手的成绩P95延迟从342ms压到58ms降幅73%显存占用从2150MB降到830MBF1保持在98.3%仅比原始模型落后0.5个百分点在业务的误差容忍范围内。8.2 踩坑清单按严重程度排序这些坑花了相当多的时间才趟平写出来希望后来人能绕开校准集选错导致PTQ精度崩坏。这个前面已经详细说了是本次项目里最严重的一个坑。吃了教训之后我的经验是上线前一定要做“校准集-部署集分布一致性检查”可以用简单的KL散度指标先量化对比两个分布的重合度再决定是否用部署集数据更新校准集。ONNX导出时忽略了动态维度。最初导出的ONNX没有声明dynamic_axes导致TensorRT编译出来的engine把batch固定成了1线上并发一上来就有大量排队差点又引发事故。剪枝后直接用旧量化参数。剪枝结束后我没有重新走量化校准流程直接载入旧的INT8配置精度掉了大约1.2%才意识到问题。正确做法是剪枝微调完成后重新收集一批校准数据从头跑一遍PTQ。TensorRT版本跨机器不兼容。开发机上编译好的engine拿到服务器上加载失败排查了大半天才定位到版本差异。现在我在CI流程里加入了一步编译engine后立刻用trtexec --loadEngine做加载验证确保产物有效。微调时把低层也解冻了。有一轮蒸馏微调我把所有层都放开更新结果模型在验证集上泛化变差了说明低层通用特征被过度扰动。改成冻结前3层之后效果立刻明显回升。8.3 我在实际项目中的排错体验整个项目做下来最大的体会是模型优化不是一个能一步到位的技术动作而是一连串依赖因果关系的决策链路。每一步优化都会改变模型的行为进而影响后面步骤的前提条件。校准集变了精度就变、剪枝变了量化配置就要重做、TensorRT编译环境变了engine就不可用——这些因果关系环环相扣任何一个环节没跟上都会打乱整体节奏。如果时间有限我会优先推荐做动态padding和PTQ量化这两件事性价比最高。如果精度预算再宽一点可以跳过剪枝和蒸馏只做量化加TensorRT也能拿到大部分的速度收益。剪枝和蒸馏的价值更多体现在显存极端受限制的场景比如边缘设备或者多模型混部。8.4 一个后续值得探索的扩展方向接下来我计划把校准集的自动化更新做成一个小工具——每次模型发版前自动从线上日志里采样一批最新请求做分布漂移检测低于阈值的直接采为新校准集高于阈值的则告警提示数据分布发生了大变化需要人工确认。这能从机制上防止PTQ部署一段时间后因为线上数据漂移导致的精度回落。类似的思路也可以延伸到蒸馏的学生模型定期重新蒸馏用新数据持续校正软标签。这个改进短期内不会带来显著的速度提升但对线上精度的长期稳定很关键——毕竟模型优化做完了真正的运维挑战才刚刚开始。
返回列表