☰
BERT GPU推理性能优化:从算子调度到融合实战
2026/10/10 4:41:28 网站建设 项目流程

BERT上线GPU后性能不合格,这件事我复盘了很久,复盘过程中收获最大的参考,是微信AI团队在PPoPP上分享的推理优化工作。先说结论:多数BERT在GPU上跑不快的根子不在显卡算力,而在算子调度和中间结果搬运——这两件事,恰好是那篇论文花大量篇幅在讲的。下面我把当时的排查链路、从论文里get到的关键思路,以及最终落到线上服务的改造步骤完整写一遍,适合正在做NLP推理、GPU利用率上不去、或者P99延迟老是超标的读者参考。

1. BERT上线GPU性能不合格,问题到底出在哪

1.1 性能验收为什么会卡在"不合格"

我遇到的项目验收标准其实不算苛刻:单卡A100上,P99延迟不超过30ms,QPS稳定在2000以上,GPU显存占用控制在16GB以内,利用率最好能过一半。结果上线压测一跑,成绩单很扎眼:P99延迟80ms,QPS只有600,GPU利用率在17%到25%之间徘徊,显存倒是没超,但波动大得离谱,像心电图的锯齿。

这个成绩单单独拎出每一项,都指向同一个诉求:GPU在大量时间处于"空转等数据"的状态。BERT这类Encoder模型跟常见的CNN图像模型不一样,它没有那种超大连续的卷积计算,而是由很多小算子拼起来的。以BERT-Base为例,12层Transformer,每一层都要做多头注意力、线性投影、LayerNorm、残差相加、GELU激活,这些算子单独拎出来都不大,尤其是短文本场景下,一批请求里平均句长可能只有三五十个token,矩阵被裁得很小,GPU的算力根本喂不满。

这里有个很反直觉的结论:显卡的FP16算力是越来越高,但BERT推理性能往往不是被算力卡住,而是被"调度开销"卡住。GPU执行一个kernel的流程是CPU端先准备好参数,再通过驱动提交队列,GPU执行完还要同步通知CPU。单个小kernel执行可能只要几微秒,但CPU提交加同步可能要几十微秒甚至更多。当模型里有一两百个这样的小kernel时,这种"启动空气"的时间占比就会迅速压过真正的计算时间。

1.2 三类典型的不合格症状,对应三种病根

我后来复盘了手头好几个项目,发现BERT GPU性能不合格基本可以分成三类。

第一类是GPU利用率低、CPU忙得团团转。这种情况最普遍,CPU端要做tokenizer、padding、特征拼接,还要一个接一个地提交kernel,Python端的开销再加一层,然后GPU一直接不到活。诊断命令一眼就能看出来:nvidia-smi里GPU利用率个位数,但top里CPU占用跑满好几核。这属于调度链路的问题,也是最值得先解决的。

第二类是GPU利用率看着还行,但延迟就是压不下来。这种一般发生在batch size偏大的时候,虽然GPU大部分时间在干活,但显存带宽被密集的中间结果搬运打满,或者padding浪费了太多算力。比如一批句子长度从10到512不等,全部pad成512,等于让GPU做了几千倍的无效计算。这属于数据形状和算子设计的问题。

第三类是显存OOM或者显存抖动剧烈。很多推理框架默认每次前向都会申请新的中间张量,动态分配器会对显存做频繁的allocate-free,碎片化以后显存看起来没超,实际分配不出连续块。这属于内存管理的问题。

我见过不少项目把所有问题都归咎于"卡不行",然后去申请A100、申请更大的机型,结果花了几倍成本,延迟没降下来多少。正确做法是先判断自己属于哪一类,再对症下药。

2. 从Profiling看真实瓶颈链路

2.1 宏观定位:NSight Systems第一锤

在动手改任何代码之前,先做profiling几乎是铁律。我自己惯用的工具是NVIDIA的Nsight Systems,它的定位是系统级分析器,负责回答"时间花在哪"。对PyTorch进程跑一次:

nsys profile --stats=true -o bert_prof python serve.py

跑完一个压测片段后,看生成的report,重点关注GPU Kernels Summary里的总时间,以及时间线里kernel之间大片大片的空白。我当时拿到的典型结论:整个压测周期内GPU实际执行kernel的时间只占18%,剩下80%以上都是CPU提交、等待同步、数据搬运的间隙。而这些间隙在时间线上一段一段显示得很清楚,kernel像孤岛一样被空白包围着。

这一步最有价值的地方在于,它先把方向钉死了:不是某个kernel写得不好,而是整个执行模型就有问题。如果一开始就直接去调单个kernel,很容易陷入局部优化,回头一看整体延迟一点没变。

补充一个细节:如果你用的是PyTorch,还可以把torch.profiler的导出结果导入NVIDIA Nsight Systems,这样能把Python侧算子和GPU kernel一一对应起来。我尤其喜欢看其中的同步等待和gap统计,它们直接反映CPU和GPU之间等待了多少时间。

2.2 微观手段:Nsight Compute看单个kernel

宏观确认是调度和搬运的问题之后,再用Nsight Compute做微观分析,确认哪些kernel本身还有救:

ncu --set full -k "regex:.*fused.*" -o bert_kernel python serve.py

Nsight Compute的前端界面会给出一堆硬件计数器数据,比如occupancy、warp stall、内存吞吐。我在BERT上看到的信息一般分两种:GEMM类kernel(MatMul、bmm)的算力利用率其实还可以,在FP16下能跑到60%以上;而LayerNorm、Softmax、Add、GELU这类纯memory-bound的kernel,内存吞吐量几乎打满,但执行时间极短,于是每次kernel launch的固定开销就被无限放大。这类小kernel是fusion的重点对象,因为它们单独再怎么优化,也快不过毫秒级的launch开销。真正的空间在于把多个kernel合并成一个,让中间结果留在片上。

Nsight Compute还有几个指标值得长期关注:SM整体利用率、kernel实际占用GPU的时间、以及launch开销统计。把这些指标组合起来,你就能知道哪些kernel值得花心思去融合,哪些kernel老老实实留给cuBLAS就好。

2.3 容易被忽略的CPU端暗伤

除了GPU自身的kernel时间,还有一块时间经常被忽略:CPU端前处理。BERT推理链路通常包含tokenizer、padding、attention mask构造、batch组装、设备拷贝,这些全都在GPU kernel统计之外,但每一毫秒都会累加到请求的端到端RT里。我当时用py-spy抓过一次CPU火焰图,发现压测时Python线程有25%的时间花在tokenizer和dict构造上,还有20%的时间被GIL的等待占掉。GPU经常空着手等CPU把batch打包好再送上来。

解决办法我给几个:tokenizer提前加载并在工作线程并行跑,输入特征用定制的numpy结构而不是Python list嵌套,batch组装完全放到独立线程,主线程只负责提交GPU任务。另外,请求进入serving框架后,避免在GPU计算路径上做加锁、日志打印、网络库解析。CPU侧的这些"杂活"清理干净之后,即便kernel不变,端到端延迟也能明显降下来。

3. 微信AI的PPoPP论文给了我哪些关键启发

3.1 我读到的论文核心判断

找优化思路的过程中,正好读到微信AI团队在PPoPP分享的Transformer/BERT推理加速工作。这篇论文让我特别对路的一点,是它完全从工业界线上服务出发,拿真实请求分布来验证方案,而不是在理想化的固定shape和理想batch上刷数字。

论文里有个让我印象极深的数据:优化前的BERT上线流程里,GPU kernel的实际执行时间只占一个请求step时间的不到三成,剩下七成多都耗在kernel launch、CPU-GPU同步、中间张量的显存写入写出上。这跟我自己profiling出来的18%到20%基本是同一个量级。论文据此提出了一个核心判断:BERT这类模型的GPU推理性能,瓶颈主要是两个,一是kernel级调度开销,二是中间结果在显存里的反复搬运。前者靠减少kernel数量解决,后者靠算子融合与内存复用解决。

这个判断本身在业界其实不算新鲜,英伟达、其他大厂都说过类似的话,但微信AI的论文把它量化到了具体模型和具体负载上,并且给了一个在工业环境里可操作的优化清单。你能看到哪些算子融合收益最大,哪些场景融合收益边际递减。这种信息比一句笼统的"要融合算子"值钱得多。

3.2 算子融合:从几十个kernel到几个kernel

论文着墨最多的就是算子融合,我把它拆成三个层面理解。

第一层是elementwise类算子的链式融合。比如标准Transformer里经常出现的LayerNorm + 残差相加 + GELU三段,单独实现是三次kernel launch,中间要写两次显存再读两次;融合成一个kernel后,中间的残差和归一化结果直接存在寄存器或shared memory里,写完就走,完全不落显存。单个融合kernel的耗时大约是原来三个kernel总和的六成,更重要的是把两次launch gap消掉了。

第二层是多头注意力(MHA)的内部融合。Q、K、V三个投影本质上可以合并成一个大GEMM;scaled dot-product attention里的QK^T、softmax、softmax(QK^T)V,在FlashAttention思路里是一个kernel完成;最后attention输出再接一层projection。论文里给出的做法和这些思路一致,但更强调工程约束:比如在GPU显存带宽有限的情况下,softmax中间的attention score矩阵不要写成全局完整矩阵,而是分块算、分块丢弃,这会显著省掉中间矩阵的带宽开销。

第三层是跨层的Launcher级融合。这个稍微激进一点,把连续几个简单操作整合进同一个CUDA Graph节点,减少CPU提交kernel的次数。CUDA Graph把大量kernel的依赖关系一次性打包提交,CPU只需要submit一次。论文的测量里,光是这一步就能把kernel launch总开销降低一个数量级。

我整理了自己项目里的融合前后对比,概念上是这样的:

算子片段优化前kernel数优化后kernel数预期效果
LayerNorm+Residual+GELU31省两次launch和两次显存往返
Multi-Head Attention8-102-3省QKV合并、分块softmax
FFN两段全连接+激活4-51-2省中间激活写读

实际收益因模型和GPU型号而异,但数量级不会骗人。

3.3 显存管理与复用策略

论文另一个重点在显存管理。BERT推理里中间张量数量巨大,循环层数一多,反复生成又销毁。如果每次都让分配器去显存里找一块连续空间,碎片化和带宽浪费都很伤。论文的方法叫memory planning或者memory arena预分配,简单说就是在一个会话开始前,把整个推理过程会用到的中间buffer大小算清楚,一次性从显存申请一块大内存,之后所有中间结果都在这个大块里按照编译期规划好的偏移量复用,不再动态申请。这跟PyTorch的caching allocator有点类似,但规划粒度更细,是按一个完整网络的前向传播来做全局规划的。

半精度也是个配套措施。BERT模型主干可以安全地切成FP16存储和计算,但LayerNorm、softmax这类对数值范围敏感的算子保留FP32。论文在这种混合精度下拿到了约两倍的吞吐提升,同时精度损失控制在可接受范围内。这个结论我后续在自己的模型上复现过,基本成立。

3.4 不要盲目追求全融合

论文里也写得挺明白:不是所有算子都该融。大GEMM(比如hidden size 768乘768的矩阵乘)本身是compute-bound,cuBLAS已经优化得很好,强行融合可能导致寄存器溢出或shared memory不足,反而更慢。融合的主要对象是那些memory-bound的小算子和elementwise算子,它们的启动开销占比高、带宽利用率低,融合收益最大。我的实践经验是:判断一个算子该不该融,先看两个数,kernel执行时间占整个step的比例,以及是不是访存密集型;如果执行时间小于几十微秒且是bandwidth-bound,基本都可以作为融合候选人。

4. 把论文思路拆成可落地的改造步骤

4.1 先换推理后端,而不是硬改PyTorch

如果你和我一样,第一反应是想直接用CUDA手写一堆融合kernel,我建议你先刹车。对于标准BERT结构,业界的成熟推理引擎已经把论文里绝大部分思路都实现了。直接切换后端,往往是最快拿到性能提升的路径。

我当时对比过的方案大概是这样:

方案部署成本灵活度实测性能提升(相对PyTorch Eager)
TorchScript编译低中1.2-1.5倍
ONNX Runtime + CUDA EP中中1.5-2倍
TensorRT(FP16+动态shape)高低2-4倍
FasterTransformer/TurboTransformers中高中3-6倍

前提是模型算子能被完整导出;如果模型里加了自定义算子,比如某种特殊的位置编码、用户自定义attention mask逻辑,导出就很容易断掉。

我当时遇到的情况是模型结构里有几个自定义OP,没法直接导出TensorRT,所以最终方案是"框架为主,手写为辅":主干走优化好的推理引擎,自定义小算子单独用融合kernel塞进去。这个组合在工程上最稳。

顺便说一句,微信相关的TurboTransformers开源项目,设计思路跟PPoPP上那篇论文是一脉相承的,里面就内置了BERT的融合算子集和显存复用机制。如果你不想碰TensorRT的配置地狱,先拿它跑通流程也不亏。

4.2 手写融合kernel的最小路线

在不得不手写的地方,我比较推荐先用Triton写,而不是直接怼CUDA C。Triton写起来像写Python,但它能直接生成比较高效的GPU代码,方便快速验证融合思路。以最常见的一段"LayerNorm + residual + bias + GELU"为例,简化的Triton片段长这样:

import triton import triton.language as tl @triton.jit def fused_ln_res_gelu(x_ptr, res_ptr, bias_ptr, out_ptr, n, eps, BLOCK_SIZE: tl.constexpr): pid = tl.program_id(0) offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offs < n x = tl.load(x_ptr + offs, mask=mask) res = tl.load(res_ptr + offs, mask=mask, other=0.) bias = tl.load(bias_ptr + offs, mask=mask, other=0.) y = x + res + bias mean = tl.sum(y, axis=0) / n var = tl.sum((y - mean) * (y - mean), axis=0) / n y = (y - mean) / tl.sqrt(var + eps) y = y * 0.5 * (1.0 + tl.math.erf(y / tl.sqrt(2.0))) tl.store(out_ptr + offs, y, mask=mask)

注意这是示意代码,真实场景还要处理多维layout和行/列规约方式,但思路就是这样:所有操作在一个kernel里完成,中间y完全不用写回显存。用Triton写完跑一次单元测试,对比PyTorch原实现的输出,误差控制在1e-4以内,就可以接进serving框架。

我踩过的坑是:不要在融合kernel里套PyTorch的autograd。线上推理不需要梯度,直接把torch.no_grad()包到整个调用路径上,并确保所有tensor都通过triton kernel的指针读写,不要让框架在中间环节偷偷插入持久化或者拷贝操作。

4.3 服务端配合优化:batching与stream

只有kernel层优化还不够,服务层如果处理粗糙,GPU依然喂不饱。

首先是动态batching。BERT推理服务收到的是不同时刻的独立请求,如果单请求进入GPU,一张A100跑一个只有几十token的矩阵,纯属浪费。正确做法是开一个请求队列,攒够一定batch大小再统一提交推理,或者在延迟容忍范围内攒固定时间窗口(比如8ms)。我上线时用动态batching把QPS直接翻了近一倍,代价是尾延迟略涨,但通过限制最大batch和超时时间可以平衡。

其次是continuous batching的思路可以借鉴,虽然它更多用在解码生成场景,但核心思想是"不让GPU等最长的那条样本";对BERT来说,可以把batch内句子按其长度分组分别提交,短句子和长句子走不同的kernel shape,减少padding浪费。

第三是CUDA stream。如果一份请求里既有前处理拷贝、又有显式计算,可以把拷贝放到独立的copy stream。论文里也提到用多stream做pipeline overlap,让GPU算第N个batch时,CPU已经在准备第N+1个batch。这一步实现起来不难,就是用PyTorch的.to(device, non_blocking=True)配合side stream,但收益很大,GPU的空隙时间会被进一步压扁。

服务端三个改动加起来,对我来说是QPS从600到1500左右的提升,之后再叠加算子融合,才站稳2400。

5. 上线验证与踩坑复盘

5.1 验收结果对比

改造完的验收数据我保留了一份,放在这供参考(A100单卡,模型BERT-Base,平均句长45 token,并发200):

指标优化前优化后
P99延迟80ms28ms
平均延迟35ms14ms
QPS6002400
GPU利用率17%-25%55%-65%
显存峰值12GB9.8GB
错误率0.1%0.1%

这个结果对应了论文里的判断:性能瓶颈在调度和搬运,只要把这两块削掉,算力红利自然就出来了。GPU利用率虽然没有到90%以上,但考虑到短文本场景本身kernel粒度小,55%-65%已经算健康。

5.2 全程踩过的坑

第一个坑是FP16溢出。把整个模型切成FP16之后,压测长句子和大batch时精度明显劣化,最后定位到attention score在FP16下数值范围不够,特别是序列长的时候,QK^T的结果绝对值偏大,softmax前的scale没做好。解决办法是把attention内部的QK^T和softmax保留FP32,只在GEMM外面用FP16。这个改动让精度恢复,性能只掉了不到5%。

第二个坑是CUDA Graph的动态shape。我本来想用cudaGraph直接把整个模型包起来,减少所有launch,但BERT推理的batch size、序列长度实时变化,每次shape一变就要重新捕获一次graph,捕获开销反而把性能吃回去了。最终的折中方案是给batch size和seq len各分几个桶,比如batch集合取1、4、8、16,seq len集合取32、64、128、256,每个桶捕获一次graph,运行时按桶命中。桶外请求回退到eager模式。

第三个坑是融合kernel的精度回归。手写kernel必须要一套强校验流程,我用的是保存一批固定输入的原始输出作为golden,每次改完kernel都跑一遍allclose,rtol设为1e-3,atol设为1e-3,不合格的kernel一律不上线。这个流程救了我至少三次。

第四个坑是共享显存碎片化。多模型实例部署在同一张卡上时,显存看起来还剩很多,但融合kernel申请连续buffer时OOM。后来给每个实例单独预分配自己的显存池,并定期在低峰期重启释放碎片。

5.3 如果重来一次,我会更早做的几件事

按我现在的经验,接手这类"GPU BERT性能不合格"项目,第一周就该完成profiling和选型评估,而不必把精力花在一个又一个低效的局部优化上。NSight跑一次半天能出结论,换推理引擎一周能完成验证,手写kernel性价比其实没那么高,只有在自定义算子成为硬瓶颈时才值得投入。

还有一点是给上游模型加一些推理友好的约束。padding策略统一、输入特征缓存、固定请求schema,这些会直接降低serving侧的复杂度,让kernel融合和batching策略更容易落地。如果让我重来,我会更早和模型训练团队对齐这些约束。

最后说点个人感受。微信AI那篇PPoPP论文最打动我的,是它把工业界真实负载里的细节摊开来了:哪条kernel路径最耗时、哪个中间张量最占带宽、哪个融合收益最大,这些信息在官方框架文档里很难找到,但做线上优化的人恰恰最需要。性能优化这种事,道理大家都懂,差的就是这些钉死的细节。我觉得做GPU推理服务的人,都可以拿这篇文章的思路当个检查清单,逐项对照自己的系统,多半能挖出意想不到的收益。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询