Transformer GPU底层优化:算子融合与KV Cache实战
1. 为什么Transformer的GPU性能瓶颈不在“模型”而在“搬运”先抛一个反直觉的结论Transformer在GPU上跑得不够快绝大多数情况下不是因为算力不够而是因为数据搬运的速度跟不上计算的速度。很多刚接触Transformer优化的同学会把注意力放在模型结构本身比如拼命调FFN的维度、换注意力头数、改层数结果发现GPU利用率还是上不去。其实当你打开NVIDIA的Nsight Systems去采样时会看到大量时间花在访存操作上——这里多等一拍那里又等一拍SM流式多处理器大部分时候在空转等待数据到位。Transformer和CNN的访存特性差异很大CNN靠卷积的权重复用能有效减少访存量而Transformer的自注意力机制需要频繁地在不同的Token之间交换信息这种模式天然就把数据搬运的短板放大了。快手那篇讲GPU底层优化的文章核心思路不是继续压榨模型的数学表达而是从GPU的执行模型出发重新审视Transformer算子在硬件上的真实行为。简单说就是把“算得慢”拆解成“数据来得慢”“指令发得慢”“小算子排队慢”三个维度分别对症下药。打个比方。你开一家餐厅厨师SM手艺再好如果传菜员显存带宽一盘一盘慢慢端顾客每个迭代照样等得抓狂。传统优化思路是换更好的厨师而底层优化的思路是重新设计后厨动线哪些菜可以一锅出算子融合哪些食材可以直接放在手边内存复用哪些工序可以并行准备并行策略从系统层面把整体吞吐拉起来。这篇文章我会从Transformer在GPU上的计算特征入手把底层优化的核心手段拆开讲清楚包括算子融合、KV Cache显存管理、Memory Bound算子的针对性优化、CUDA Graph的启动开销消除以及多卡场景下的负载均衡。这些都是快手那类工业级优化里真正用到的东西不是什么花架子。2. Transformer在GPU上到底“慢”在哪从计算特征说起2.1 算得快还是搬得快Transformer是Memory Bound还是Compute Bound判断一个模型在GPU上的表现首先要搞清楚它是Compute Bound算力受限还是Memory Bound带宽受限。两者的优化方向截然不同Compute Bound要靠提高计算密度让SM始终有活干Memory Bound则要想办法减少数据搬运量把访存路径上的瓶颈打通。Transformer的算子在Bound类型上是分裂的这是它特有的麻烦。Matmul矩阵乘法类算子比如QKV投影、FFN的两个线性层、注意力分数的变换这些属于典型的Compute Bound。尤其是当Batch Size足够大、矩阵维度足够宽时Tensor Core可以发挥出很高的峰值算力。注意力部分包括Softmax、Mask、Attention Score的计算以及最终的加权求和在长序列场景下很容易变成Memory Bound。因为注意力分数矩阵的尺寸是batch_size × num_heads × seq_len × seq_len随序列长度平方增长。这个中间矩阵如果写回显存再读出来带宽消耗极其惊人。更麻烦的是Transformer里小算子特别多。比如Residual Add、LayerNorm、激活函数、Mask操作、Tensor的维度变换Reshape/Transpose等等单个算子计算量不大但每个都要从显存读一遍、写一遍。这就像每道菜都单独叫一个传菜员后厨人再多也不够用。2.2 一个小实验为什么加个Transpose就能掉30%性能这里分享一个我实测过的案例。某次做GPT风格解码优化时只是把Attention里的Score矩阵在计算过程中顺手做了一次Transpose让布局更符合后续算子的要求结果端到端性能掉了接近30%。原因不复杂。Transpose本身几乎不消耗算力但它把整个矩阵从显存读出来、重新排列、再写回去这一来一回消耗的带宽和时间远超想象。在Memory Bound的阶段一个多余的数据排布转换就是灾难。这也是为什么真正做底层优化的人会对Tensor的内存布局Layout极度敏感因为布局错了整个计算链路上的每一个算子都要跟着做额外的读写。2.3 GPU执行模型角度SM空转的三大原因从底层看SM空转通常有三个原因访存延迟未隐藏。GPU靠大量线程并行来掩盖访存延迟如果并行度不够或者每个线程访存的局部性很差SM就只能等数据回来再继续算。算子粒度太小。一个小算子启动后只运行几微秒就结束了大量时间浪费在Kernel Launch内核启动的开销和线程调度上。依赖链太长。Transformer层内存在天然的串行依赖比如LayerNorm需要等上一层的输出注意力Softmax需要等完整的Score矩阵。依赖链不打破就算有再多的SM也有一部分在空转等待。快手那篇文章里提到的底层优化核心就是围绕这三点做文章要么让算子更粗更快要么让内存布局更合理要么从算法层面打破依赖链。3. 算子融合把后厨的传菜员减掉一半3.1 为什么融合能提速从“读一遍写一遍”到“原地算完”算子融合是现在Transformer优化的基本操作也是效果最立竿见影的手段。它的思路很简单把多个相邻的算子合并成一个Kernel避免中间结果反复读写显存。以LayerNorm为例。常规实现分三步计算均值mean和方差var。对每个元素做归一化(x - mean) / sqrt(var eps)。乘以缩放参数gamma和偏移参数beta加上残差连接residual。如果每个步骤单独一个Kernel那么中间结果至少要完整地写回显存一次、再读出来一次。而融合成一个Kernel后整块数据在SM内部的寄存器或共享内存里就完成了所有操作对外只读一次、写一次效果立竿见影。3.2 融合的边界在哪里哪些能融哪些不能硬融不过算子融合也不是无脑融。需要遵守几条原则访存密集型算子优先融合。像LayerNorm、Softmax、激活函数这种几乎不消耗算力、纯粹在搬运数据的算子融合收益最大。避免大中间结果物化。比如Attention的Score矩阵如果能在计算Softmax之前就做Mask并且以分块的方式在片上完成就不需要把完整的Score矩阵写回显存。融合后共享内存别爆掉。GPU的共享内存是有限的比如A100单个SM是164KB如果融合的算子太多导致共享内存占用超标反而会降低SM上同时运行的Block数量得不偿失。3.3 实操FlashAttention式的分块融合思路FlashAttention是注意力算子融合的典型代表它的核心思想就是不要物化完整的注意力分数矩阵而是分块计算并在计算过程中同步更新输出。简化版流程把Q、K、V都切成Block。每次取一个Q Block和K Block算局部Score。在片上用Online Softmax的增量公式更新概率和输出。所有Block遍历完输出直接就是完整结果。这样整个注意力的计算过程中需要写回显存的只有最终的输出而不是巨大的Score矩阵。长序列场景下这种优化能把注意力部分的显存占用从O(n²)降到O(n)速度提升通常在一倍以上。提示如果你用PyTorch可以优先尝试torch.nn.functional.scaled_dot_product_attention它在Hopper架构上会自动选择FlashAttention或类似的高效实现。但如果你要部署到自己的推理框架里理解分块融合的原理依然是必须的因为SDPA的自动调度未必适配你的内存布局。4. KV Cache解码阶段的隐形吞吐杀手4.1 为什么自回归解码会变成“算力过剩、带宽吃紧”Transformer在做生成解码时每个新Token都要和之前所有Token的Key、Value做注意力计算。如果不做任何优化每个Step都要把历史Token的K、V重新算一遍这显然是巨大的浪费。所以常规做法是把历史的K、V缓存起来这就叫KV Cache。但KV Cache的引入带来一个系统级麻烦随着生成步数增加缓存的数据量线性增长。以LLaMA-7B为例单序列的KV Cache大小大约是2K和V × num_layers × num_heads × head_dim × seq_len × 2字节算下来生成2048个Token时缓存接近100MB。这意味着解码阶段的每一次前向推理都要把和模型大小同量级的KV Cache读一遍。模型参数也占带宽缓存也占带宽两者叠加后Memory Bound的特征被进一步放大。这就是为什么在解码阶段GPU算力往往过剩但Token生成速度上不去——瓶颈完全在带宽。4.2 PagedAttention与KV Cache的显存碎片化缓存大了之后显存分配变得棘手。不同序列长度不同动态分配KVCache时容易产生碎片化浪费大量显存。vLLM提出的PagedAttention思路借鉴了操作系统里的虚拟内存分页把KV Cache切成固定大小的Block用页表来管理序列到物理块的映射。这样做的好处按需分配不需要为每个序列预留最大长度的连续显存。吞吐提升可以同时容纳更多序列更大的Batch摊薄模型参数的访存开销。内存碎片率大幅下降。我个人在部署推理服务时显存利用率和吞吐量对比过用PagedAttention之前一个8卡A100服务能跑的并发是50左右换成支持PagedAttention的推理框架后同样的硬件并发能到200以上。这个差距不是模型优化能追回来的。4.3 降低KV Cache带宽压力的几个实战手段除了PagedAttention管理显存还有几个方向的工程实践值得记录GQA / MQA分组查询注意力 / 多查询注意力让多个Query头共享同一组Key、Value。这样KV Cache的体量直接缩小数倍带宽压力同步下降。LLaMA-2 70B、Mistral等模型都已经在用这种方式换推理速度。KV Cache量化把缓存从FP16压到INT8甚至INT4。实验显示适度量化KV Cache可以在几乎不掉点的情况下把缓存的访存量再减半。缓存滑动窗口对超长序列做流式处理时没必要保留全部历史KV只需要保留窗口内的部分。这在流式交互场景里非常实用。5. CUDA Graph与内存池消除启动开销和反复分配的隐藏成本5.1 Kernel Launch开销小算子的“慢性毒药”传统PyTorch执行模型里每一个算子都是一次Kernel Launch每次Launch都要经过CPU发指令、GPU接收、调度执行的过程。单个Kernel的Launch开销在几微秒到十几微秒听起来不贵但Transformer的一次前向推理里有几百个算子累计起来就是几百微秒到毫秒级的开销。在小Batch或实时性要求高的场景里这部分的占比相当可观。我之前在一个实时交互场景里实测短序列下Kernel Launch开销能占到总时延的30%以上。尽量缩减Launch次数就变成了一个必须解决的问题。5.2 CUDA Graph解决的是什么问题CUDA Graph的思路是把一连串的Kernel Launches提前录制下来在GPU上形成一个完整的依赖图运行时一次性提交执行。这大大削减了CPU和GPU之间的交互次数。实操中需要注意图捕获Capture阶段不做动态内存分配否则捕获会失败。标准做法是提前分配好内存池在捕获期间复用现有显存。图的输入输出需要用固定地址的缓冲区比如torch.cuda.graphs里的TensorPool机制。动态Shape场景要小心。CUDA Graph要求捕获时的Shape和实际运行时一致如果序列长度变化要么按最大长度Padding要么为不同Shape分别捕获Graph。FlashAttention这类算子在图捕获模式下通常能正常工作但融合算子如果内部有依赖输入Shape的动态选择逻辑就得留意会不会触碰到代码里“不均匀”的分支。5.3 显存分配也是隐形开销从torch.cuda.caching_allocator说起PyTorch默认的显存分配器已经做了缓存不会每次分配都向驱动申请显存但它依然会在每次分配/释放时加锁、搜索空闲块在多线程并发场景下锁竞争非常明显。优化手段包括使用更大粒度的内存池减少零散分配。在服务框架层自己做显存复用比如把不同序列的中间结果分配到同一个缓冲区。推理引擎TensorRT-LLM、vLLM等内部的Allocator通常已做了优化自己写框架时最容易忽略的就是这一层。6. 并行策略单卡优化之外多卡如何分摊Transformer的算力和带宽6.1 Data Parallel与Tensor Parallel的边界单卡优化到极限后下一步就是多卡并行。遇到超大模型几十B甚至上百B参数时单卡显存放不下必须做模型并行。最常用的是Tensor Parallelism把权重切到多张卡上每张卡算一部分通过AllReduce汇总。Tensor Parallel的核心问题是通信开销切分得越碎通信占比越高。实操中需要找trade-off点。以LLaMA-65B为例常见的8卡TP配置下单次通信量已经接近单层激活值的大小如果切到16卡通信对端到端吞吐的影响会明显增加。这也是为什么张量并行通常控制在8卡以内更大的规模要用流水线并行或DPTP的组合。6.2 Pipeline Parallelism的切分艺术Pipeline Parallelism把模型的Layer切到多张卡上每张卡负责一部分层。好处是通信量比TP小但会引入Pipeline Bubble流水线气泡即某些卡在等待上游结果时空闲。实践中建议尽量保证各Stage的计算量均衡避免某张卡成为瓶颈。Stage间只传递激活值中间结果不要带太多额外信息。用Virtual Pipeline或Interleaved策略把Bubble进一步摊薄。6.3 通信与计算重叠AllReduce不一定要等多卡训练时每次AllReduce都要等所有卡算完才能开始通信时间常常暴露在关键路径里。优化的办法是通信与计算重叠Overlap在计算下一层的时候把上一层的梯度或激活异步通信出去。实操中常用的技巧是把一个大的AllReduce拆成多个小的AllReduce让后一个计算块和前一block的通信同时进行。对于自注意力里的AllReduce还可以在层内做额外的切分把一个Layer内部的通信任务拆细让通信和Matmul尽可能重叠。7. 工业级实战中的性能排查链路从Profiling到定位瓶颈7.1 第一步先量化别瞎猜每次有人跟我说“我模型跑得慢”我的第一反应都是先给我看Profile数据别猜。工具无非这几个Nsight Systems看全局的Timeline、CPU/GPU利用率、Kernel分布。Nsight Compute看单个Kernel的内部指标比如SM占用率、访存带宽利用率、Warp Stall原因。PyTorch Profiler快速区分算子的CPU耗时和GPU耗时。如果你发现某个Kernel把GPU利用率拉满了但整体吞吐就是不涨那大概率是Memory Bound如果发现大量细碎的小Kernel那就是Launch开销和融合度的问题。7.2 一张实用的瓶颈定位表现象可能瓶颈优先排查方向GPU利用率长期低于50%访存延迟/数据依赖检查是否Memory Bound、是否Kernel太碎某个大Kernel耗时异常高Compute Bound看Tensor Core是否启用、精度是否合适小Kernel数量极多启动开销考虑CUDA Graph、算子融合多卡吞吐不随卡数增长通信瓶颈检查通信占比、通信与计算是否重叠显存占用超过预期中间结果物化/KVCache检查是否有大Tensor写回显存7.3 一次真实调优的完整链路复盘这里以一个7B模型在A100上的推理优化为例完整过程供参考基线纯PyTorch FP16推理Batch1生成速度约15 tokens/s。第一步加算子融合。用自定义的Fused LayerNorm Residual生成速度到20 tokens/s。瓶颈从访存密集算子开始缓解。第二步换FlashAttention。序列长度2560生成速度到28 tokens/s。注意此时更多的改善来自于显存占用下降Batch可以往上提了。第三步上CUDA Graph。生成速度到33 tokens/s。这里有个坑需要把输入输出固定到固定地址的Tensor上。第四步KV Cache优化。Int8量化KV Cache之后Batch从8提到16整体吞吐翻倍。第五步上vLLM这类推理框架PagedAttention 连续Batch动态调度单卡并发吞吐再次翻倍。每一步改完都能看到明确的数据变化这就是底层优化的核心魅力不需要玄学一切以数据说话。8. 一些踩过的坑和个人的经验判断8.1 算子融合不是越狠越好有一段时间我迷信完全融合使劲把多个Layer拼进一个Kernel结果某些融合版本在各种边界Shape上代码崩了调试成本极高最后回退到部分融合。还是那句话优先融合Memory Bound算子Compute Bound算子保持独立。融合的收益要拿Profile说话别为了“看起来酷”去融。8.2 先对齐数据布局再谈其他优化很多性能问题根本不是算法问题而是张量布局和算子期望的布局不匹配带来的隐式转置和拷贝。在动手写任何自定义Kernel之前先花半小时把整个计算链路里每个算子的输入输出Layout捋一遍。真实项目中这一步省下的时间往往是最大的。8.3 版本和硬件差异比想象中更大Tensor Core在不同GPU架构下的行为差异很大V100的FP16 Tensor Core远弱于A100A100的BF16支持又比FP16更稳Hopper架构则引入了FP8。你在网上看到的优化经验一定要对着自己的卡做验证别直接抄。8.4 推理框架能解决问题就不要自己造轮子如果你只是要把一个Transformer模型部署上线优先用FastTransformer、TensorRT-LLM、vLLM这些经过工业级验证的框架。自己造轮子学习可以生产环境没必要。底层优化的知识让你能理解这些框架的原理、知道怎么调参、遇到问题时能定位这才是它最大的价值。我在实际项目中反复体会到GPU底层优化不是“锦上添花”而是Transformer规模化落地的必答题。模型结构决定计算量的上限但GPU底层优化的水平决定了这个上限到底有多少能变成真实的吞吐数字。如果你刚起步建议先从Profiling开始建立“花时间在哪里”的敏感度再一步步上手算子融合、CUDA Graph、缓存优化这些手段。用数据驱动优化永远比感觉可靠得多。