Gated DeltaNet GPU并行实现:拆解线性注意力循环的两种高效方案
在GPU上并行实现Gated DeltaNet这类Linear Attention模型是我最近大半年里反复折腾的一件事。起因很直接模型的前向推理如果用朴素的串行循环去更新状态在长上下文比如十几万token下慢得没法看而几乎所有能把复杂度从O(N²)压到O(N)的线性注意力方案都依赖一个隐藏在公式里的循环状态。“循环是什么时候引入的、能不能被并行展开”就成了GPU实现绕不开的问题。这篇就围绕GDNGated DeltaNet来聊顺带把Linear Attention常用的两类并行实现思路一起讲清楚——Parallel Scan和Chunked递推。适合正在给线性注意力/线性RNN类模型写推理kernel的朋友也适合刚接触这些架构、想知道“为什么代码里没有for循环”的读者。1. 为什么说Linear Attention的GPU实现瓶颈从来不在矩阵乘1.1 GDN为什么会出现在LLM里先对齐一下背景。Softmax注意力的问题大家都清楚注意力矩阵是N×N的KV cache要存N个key和value推理时每生成一个token都要和前面所有token做一次点积长上下文场景下内存和耗时都是灾难。Linear Attention的思路是把注意力核函数做近似让状态可以累积成一个固定大小的矩阵或向量把O(N²)的复杂度压到O(N)。GDN属于这条路线里比较能打的一支它把Delta规则Delta Rule和门控机制结合起来用一个固定维度d的状态矩阵维护“当前上下文里该记住的信息”逐token更新。近一两年不少开源大模型把这类结构放进混合架构一部分层用softmax注意力一部分层用线性注意力原因就是它能在长上下文下把KV cache压到很小同时保持不错的表达力。在GPU上实现GDN矩阵乘本身不是难点——d维的key、query、value之间做乘加GPU天生擅长。真正棘手的是那个“逐token更新状态”的循环。1.2 循环依赖才是真正的对手GPU的并行度来自大量线程同时干活。一个kernel能跑多快很大程度上取决于每个线程算完一个结果后需不需要等别的线程的结果。GDN的状态更新长这样下一个时间步的状态依赖当前时间步的状态。这个依赖关系把整条序列串成了一条链——理论上第N个token的状态必须等前面N-1个token的状态算完才能算出。这就是典型的循环依赖sequential dependency。如果直接按这个依赖写代码每个时间步做一遍矩阵乘加那GPU上万个线程只有一小部分在干活剩下都在空转等数据。序列越长浪费越严重。很多人说“Linear Attention在GPU上不好实现”其实所有难点都集中在这个循环上怎么在保证结果一致的前提下把这条链拆开让尽可能多的计算并行发生。1.3 一个朴素的循环实现有多浪费我最初验证GDN正确性时直接在PyTorch里用for循环逐token更新状态。代码没有问题但性能惨不忍睹序列长度N8192d64单卡A100一个batch逐token循环的前向耗时是FlashAttention的20倍以上。问题不是算得慢是每次循环只调用一个小矩阵乘法kernel启动开销、线程利用率全部拉胯。GPU的利用率经常在5%以下大部分时间在等循环结束。这个例子说明GDN在GPU上能不能跑起来取决于能不能把序列维度上的循环并行化而不是取决于矩阵乘本身快不快。下面这两条路线都是冲着解决这个循环依赖去的。2. 抓住GDN状态转移的“线性骨架”并行化就开始了2.1 Delta规则更新不是一个外积而是一个仿射映射GDN的核心状态更新可以写成这样一个雏形S_t S_{t-1} - β_t (S_{t-1} k_t - v_t) k_t^T看着绕拆开就清楚了。这里S是d×d的关联记忆矩阵k_t是当前token的keyv_t是valueβ_t是学习率一个由输入计算出来的标量。把括号展开S_t S_{t-1} - β_t (S_{t-1} k_t) k_t^T β_t v_t k_t^T注意第一项里k_t k_t^T是一个d×d的秩1矩阵。所以S_t (I - β_t k_t k_t^T) S_{t-1} β_t v_t k_t^T这是一个标准的仿射映射Affine Mapping新状态 矩阵 × 旧状态 常数矩阵。这里的“矩阵”是(I - β_t k_t k_t^T)常数矩阵是β_t v_t k_t^T。这句话是理解后面所有并行的钥匙。很多人以为Delta规则更新只是“往状态里加一个外积”其实它同时包含了一层“擦除旧信息”的操作——(I - β k k^T)在k方向压制旧的关联再写入新的key-value关联。这个擦除项让DeltaNet区别于普通线性注意力但也让状态转移变成了一个线性变换这反而给并行化铺了路。2.2 门控的角色遗忘系数与“局部学习率”Gated DeltaNet在Delta规则之上引入了门控系数。门控通常由当前输入通过sigmoid算出来取值范围在(0,1)之间。它的作用是控制“旧状态保留多少、新信息写入多少”。加进门控后状态更新的整体结构仍是仿射映射只是系数和常数项都对应缩放。如果门控是逐维向量或对角矩阵那仿射变换里的系数矩阵会多一个对角因子如果门控是标量就直接吸收进β_t。无论哪种细节最终都能归纳成S_t A_t S_{t-1} B_t其中A_t由当前token的key和门控决定B_t由value、key和门控决定。从计算图的角度看整个时间序列的状态演化就是T个仿射映射依次作用在初始状态上。2.3 仿射转移能组合scan的数学基础为什么仿射结构重要因为两个仿射映射可以组合成新的仿射映射。假设第1到第t步的转移函数分别是f_1, f_2, ..., f_t每个都是f_i(x) A_i x B_i。它们复合的结果仍然是一个仿射映射f_t ∘ f_{t-1} ∘ ... ∘ f_1 (x) A_{t:1} x B_{t:1}其中A_{t:1} A_t A_{t-1} ... A_1B_{t:1} A_t A_{t-1} ... A_2 B_1 ... A_t B_{t-1} B_t。这个性质意味着任意一段连续区间上所有状态转移可以“折叠”成一个等价的仿射映射。而折叠操作本身满足结合律——先折叠[1,2]再折叠[3,4]和先折叠[2,3]再折叠[1,4]结果一样。结合律正是GPU并行扫描算法能成立的前提。这里要强调一点GPU上常用的并行扫描Parallel Scan针对的是满足结合律的二元运算。不是所有循环都能并行化但所有能写成仿射映射的循环都能用scan并行。GDN恰好属于这一类这也是它比某些更复杂的非线性RNN好实现的原因。3. 两条主流实现路线Parallel Scan与Chunked递推3.1 Parallel Scan全局并行但通信密集理解了仿射折叠Parallel Scan的思路就很好说了。它把整个序列的T个仿射映射均匀分到多个线程/线程块上然后分log_2 T轮做折叠第一轮相邻两个区间各自折叠成1个仿射映射得到T/2个“区间等效映射”。第二轮相邻折半合并得到T/4个映射。依此类推T轮后得到全局等效映射再作用到初始状态上一次性算出所有时间步的状态。所有折叠操作可以并行执行每个线程只负责自己区间内的矩阵乘复杂度从O(T)降到O(log T)轮。但Parallel Scan不是免费的午餐代价集中在通信和中间存储每轮折叠都需要把相邻区间的A、B矩阵交换给对应的线程这在GPU上是全局内存读写或warp shuffle。A矩阵是d×d的中间结果要存O(T d²)数据如果所有轮次都保留下来内存占用很可观。实际工程里通常只保留每轮的折叠结果甚至用“两阶段scan”来减少存储。当d变大时A矩阵本身的计算和通信成本按d²增长scan部分会越来越重。我在实践中对Parallel Scan的主要印象是它适合d较小比如64或128、序列不长几千到几万的场景。d一旦到512scan里的矩阵乘规模和通信量会迅速盖过收益。3.2 Chunked方式把序列切成块块内做矩阵乘Chunked方式是对Parallel Scan的一个工程化改良思路极其直观把序列切成大小固定的块chunk每个块内部的token并行处理块与块之间串行传递状态。具体拆解一下对第j个块先并行算出该块内所有token对状态的“局部贡献”和“局部转移”比如块内第i个token的key k_i、value v_i组合成的秩1矩阵k_i k_i^T以及对应的常数项v_i k_i^T。这些局部量的计算是纯粹的矩阵乘和块内其它token没有依赖可以完全并行。然后从第一个块开始依次把块的局部转移乘到累积状态S上再把贡献加进去得到块末状态作为下一个块的初始状态。块内并行的好处是大部分计算key和value的组合、输出投影都转成了大矩阵乘法能充分用上Tensor Core。块间虽然还是串行递推但串行轮数从T降到了T/BB是块大小通信成本大幅下降。前文提到的备受欢迎的开源库FLAFlash Linear Attention对多个线性注意力模型包括DeltaNet、Gated DeltaNet、GLA、Based等都采用/支持这种chunked框架块大小通常取64或128。3.3 双并行/Chunk-wiseGDN自己的工程选择如果只看Chunked还有个问题没解决序列很长时块的个数T/B还是很大比如N131072B64块的个数是2048。如果每个块由一个线程块负责2048个线程块在GPU上虽然能同时跑一部分但块间状态传递还是需要一个全局的串行依赖。GDN论文里明确提出了两档并行Block-parallel模式每个线程块负责一个chunk块内并行块间需要每次同步状态。Chunk-wise模式一个线程块负责多个chunk把块间的部分状态递推留在线程块内部完成减少跨线程块同步的开销。实际实现里常见做法是“双并行”一段序列内部做块间递推同时多个线程块各自处理不同的序列段等各自算完后再做一次跨段折叠。这相当于把Parallel Scan和Chunked结合块内用矩阵乘块间用scan。我自己的体会是GDN论文的“chunk-wise”提法后来被很多库吸收本质就是用可控的chunk大小在并行度和内存占用之间做权衡chunk越大块内并行度越高但中间状态缓存也越大chunk越小块间串行轮数越少但块内矩阵乘的规模太小Tensor Core喂不饱。3.4 路线选择对照表维度Parallel ScanChunked / Chunk-wise核心思想把仿射映射折叠成结合律的scan块内矩阵乘 块间状态递推并行粒度全序列并行折叠每轮reduce块内并行块间串行中间存储O(T d²)每轮总量偏大O(B d²)可控制适合场景d小、序列短、需求简单任意d长序列工程化更稳Tensor Core利用率低scan中运算不规则高核心计算是矩阵乘上手难度数学简单kernel细节多工程结构复杂但调优空间大这个表是我个人的经验总结不是绝对的。两条路线在具体模型上经常混用很多高性能实现都是先chunked把输入切成块再在每个块内部或块间用scan做状态合并。4. 从公式到CUDA Kernel布局、块大小与Warp分工4.1 内存布局和转置被低估的头号性能杀手写kernel之前内存布局就得想清楚否则后面全白做。GDN的输入通常按(batch, head, seq_len, d)排布但矩阵乘最舒服的布局是让d维度连续。于是麻烦来了key、value参与状态更新时需要对每个时间步做k_i k_i^T这样的外积天然希望按token的d维连续访问。query输出时需要把当前状态S和query做乘加状态是d×d矩阵希望d维连续。如果直接使用原始布局核心里到处是“取一列”的操作一次转置就能吃掉30%以上的性能。我建议的实现方式是在kernel开始时做一次共同的布局转换把key、value转成(d, seq_len)的列主序布局让后续的访存全部按d维连续。不要在每个小算子后面做转置那样转置开销会重复计算。HopperH100上有TMATensor Memory Accelerator可以异步搬运和转置数据能把这步做到几乎不占额外时间AmpereA100上没有TMA就得用共享内存手动转置注意避免bank conflict。4.2 块大小B和状态维度d怎么选这两个参数是kernel性能的核心旋钮d状态维度GDN常见配置是64或128。d越小状态矩阵越轻scan的通信开销越小但表达力下降d越大表达力强但矩阵乘和scan的成本都按d²涨。B块大小我测试下来B64和B128是最稳的区间。B太小块内并行度不够B太大块间串行轮数少但每块的中间缓存S_localB×d×d量级会爆共享内存。一个参考原则让块内矩阵乘的尺寸至少达到Tensor Core单次MMA操作能高效处理的范围通常是16×16×16或更大同时让B×d不超过共享内存可用容量的四分之一。比如d128时B64的局部状态缓存是64×128×128×4字节假设fp32中间量 4MB这已经超过A100的共享内存上限了所以实际实现中中间状态通常存全局内存、用L2缓存兜着共享内存只放当前正在计算的一小片。4.3 Warp分工与共享内存分配块内并行不是简单地让所有线程算同一个矩阵。我发现最有效的做法是warp specialization按角色分工一部分warp专门做矩阵乘对block内B个token的key/value做外积累加或对输出做投影这部分用warpgroup级MMA指令吃满Tensor Core。一个或少数warp专门做状态递推把当前块的局部转移乘到累积状态S上这步是串行的但已经折叠到T/B轮开销可控。共享内存里通常放两块一块存当前块的key/value切片一块存正在累积的状态S。双缓冲可以让矩阵乘和状态递推重叠。我在A100上实测下来的典型分配一个块128个线程4个warp做矩阵乘1个warp做状态递推共享内存留16KB给数据切片、16KB给状态剩余留给算子中间量。这个比例在不同d下要微调但基本逻辑不变把最重的矩阵乘分给大部分线程把串行递推压缩到少数warp避免所有线程都卡在状态更新上。4.4 数值精度fp32累积值得较真GDN的状态更新和softmax attention有个本质区别softmax注意力每次都是独立归一化的误差不太会累积GDN的状态是持续复用的一个时间步的小误差会被后续token一直继承。我踩过的坑之一是用fp16算状态转移时beta项和门控的乘积在长序列下会产生显著的累积漂移最后导致loss出现莫名其妙的抖动。后来仔细排查不是因为算法变了而是精度不够。工程上的折中方案状态转移矩阵A_t和常数项B_t用fp32计算和存储因为它们是“每次都要乘进旧状态”的东西误差按指数级传播。矩阵乘的输入输出可以用bf16/tf32利用Tensor Core加速。如果一定要在fp16下走至少保证状态S以fp32累积在每次写回前再截断到fp16。H100上可以用warpgroup的混合精度特性让累加器保持fp32。5. 实测对比与踩坑心得5.1 三类实现的性能对照我拿一个典型的GDN层d64head8batch1序列长度N16384在A100上做过对照。这个配置接近一些开源混合架构里的实际层设置实现方式前向耗时(ms)相对baseline加速备注PyTorch逐token for循环约95ms1×明显有kernel启动开销Triton手写Parallel Scan版约12ms约8×scan的矩阵组合阶段仍有空转Triton/CUDA混合Chunked版约5ms约19×块内矩阵乘、块间递推接近理论峰值FLA现成实现约4.5ms约21×各项优化完整可当基准这里要说明以上数字只是我单卡单次运行的参考值不同驱动版本、cuBLAS版本、网络结构下会有浮动。但趋势是一致的——只要把序列维度的循环拆到块内去收益是数量级的。序列继续拉长到N65536以上时Parallel Scan版的劣势会更明显因为中间状态的通信量线性上涨Chunked版反而更稳因为块大小固定块间轮数增加但每轮的成本可控。5.2 几个容易被忽略的坑第一个坑scan的“可结合性”被破坏。我很早前手动实现parallel scan时把A矩阵折叠写成A_2A_1但折叠常数的顺序写反了。注意组合时是“后发生的映射先作用在旧状态上”顺序错了结果完全不对而且错得还很隐蔽——短序列下loss勉强能收敛序列一长就彻底爆炸。建议先用d2的随机数据做数值梯度比对确认kernel结果和naive循环完全一致再继续。第二个坑H100上用TMA时内存必须满足一定的对齐要求如果输入是从某个模型输出的非连续view传入kernelTMA会静默出错或性能暴跌。在PyTorch侧先做.contiguous()或者kernel内部先做一次detect-remap能省很多排查时间。第三个坑Gate和β的初始化尺度。GDN里的门控是sigmoid输出初始bias如果设得不好会让状态长期接近1.0也就是“永远不遗忘”导致后面序列越长累积误差越大。训练和推理两侧最好都确认门控的初始分布是偏向遗忘的比如bias往负方向初始化kernel实现上也要避免对门控做不必要的数值规范化因为它本身就是范围受限的。第四个坑把scan中间结果全量缓存到全局内存。做Parallel Scan时图方便把每轮的A、B矩阵都保存下来内存占用翻了几倍而且最后几轮访问基本不命中L2性能反而下降。更好的是两阶段scan第一遍只记录块级摘要第二遍用摘要回推最终状态把中间结果压缩一个量级。5.3 一点个人建议如果你只是需要让GDN跑起来而不想从头写kernel优先用现成的线性注意力库比如FLA这类成熟实现把精力花在理解chunk大小和数值精度这两个旋钮上。如果你是想学习GPU并行实现的思路我建议从“写一个naive循环baseline”开始然后按这两步递进先实现一个能跑对的Parallel Scan再改成Chunked。两版一对比你对“为什么GPU上矩阵乘比scan香”的理解比看十篇论文都深刻。GDN的并行实现数学上不难难的是每一步都在和硬件的脾气较劲。把循环拆开、把矩阵乘喂满、把精度稳住——这三件事做到位一个可用的GPU实现就立住了一大半。