最近被问得最多的一个问题:GDN(Gated DeltaNet)这类Linear Attention,在GPU上到底怎么并行实现?老实说,这个问题比“怎么调FlashAttention”要复杂一个档次。标准Attention的并行套路已经被FlashAttention讲得很透了,但像GDN、Mamba、RWKV这类递归式的序列模型,每个位置都在更新一个隐状态,天然是串行的,你没法简单地把一个长序列扔进大矩阵乘法里完事。可偏偏这类架构又特别看重长序列,L一上到几十万,GPU上的实现效率就直接决定这个模型能不能落地。
这篇文章我打算从原理到实操,把GDN和其他Linear Attention的GPU并行实现讲清楚。内容会覆盖状态递推的数学结构、scan和chunk这两种核心并行策略、具体的kernel实现思路,以及我在实际调优时踩过的坑。适合正在研究线性注意力实现、想手写算子、或者准备在项目里用上这类模型的工程师和研究者。
1. 要并行的是一个什么样的递推
1.1 为什么线性注意力会有状态
标准Attention里,每个token会跟前面所有token做一次点积注意力,矩阵形状是L×L,复杂度O(L²),长序列根本扛不住。Linear Attention的思路是抛弃显式的注意力矩阵,把历史信息压缩进一个固定大小的状态矩阵S里。以最简单的linear attention为例:
S_t = S_{t-1} + k_t v_t^T
输出时用查询向量去读状态:o_t = q_t^T S_t。这个状态S是D×D的矩阵,D是特征维度。每来一个新token,就往状态里加一项k_t v_t^T,于是复杂度从O(L²)降到了O(LD²),当D远远小于L时,这是本质性的改善。
注意看这个递推:它不需要看前L个token,只依赖t-1时刻的状态。这就是递归序列模型,但同时也意味着GPU上不好并行了——因为你看起来必须一个token一个token地往后算。
1.2 GDN的更新规则与可并行线索
GDN的全称是Gated DeltaNet,它把DeltaNet的delta更新规则和门控机制结合到了一起。原版DeltaNet的状态更新可以写成:
S_t = A_t S_{t-1} + B_t
其中 A_t = I - β_t k_t k_t^T,B_t = β_t k_t v_t^T。这里的β_t是一个0到1之间的学习率标量,sigmoid输出。也就是说,每次更新状态时,不是简单地把新内容加进去,而是先把旧状态中与k_t方向相关的部分擦掉一部分,再写入v_t的信息。这比“往记忆里硬加一条”要聪明,可以在一定程度上避免key相似导致的记忆污染。
GDN加入门控α_t后,更新变成:
S_t = (1 - α_t β_t k_t k_t^T) S_{t-1} + α_t β_t k_t v_t^T
从形式上看,它依然是一个线性递推:
S_t = A_t S_{t-1} + B_t
关键区别在于,这里的A_t不是一个简单的标量衰减系数,而是一个D×D的矩阵,具体是对角矩阵(单位阵)加上一个秩一修正项。这就给并行化带来了麻烦,但也留下了线索:无论A长什么样,只要递推是线性的,那整个序列的信息迁移就可以被“打包”成可以结合的算子,进而用并行扫描来处理。
1.3 和其他Linear Attention的亲缘关系
这里值得多说一句,因为游客常把GDN、Mamba、RWKV、Gated Linear Attention(GLA)当成完全不同的东西,其实它们在状态递推层面上是同一族。
Mamba的离散SSM更新是:
S_t = diag(A) S_{t-1} + B_t u_t
衰减矩阵是对角的,所以每个维度独立。
RWKV的时间混合可以改写为指数衰减的线性attention,状态更新里有一个标量衰减项。
GLA(Gated Linear Attention)的更新是:
S_t = diag(α_t) S_{t-1} + k_t v_t^T
门控α_t是一个向量,逐元素控制遗忘。
看,这些本质上都是 S_t = A_t S_{t-1} + B_t。区别只在A_t的结构:标量、对角矩阵、还是“单位阵加秩一修正”。而GPU并行化的核心,恰恰就是围绕这个A_t展开。
2. GPU并行化的两根支柱:scan与chunk
2.1 从前缀和说起:把递推变成可结合的算子
你如果写过CUDA,应该听过并行扫描(parallel scan)或前缀和(prefix sum)。它的经典场景是:给定数组x_1, x_2, ..., x_L,计算每个位置的前缀和 y_t = x_1 + ... + x_t。表面上这也是串行递推,但因为加法满足结合律——(a+b)+c = a+(b+c)——所以可以用树形归约并行:先两两相加,再四四相加,log L轮就完事。
线性递推 S_t = A_t S_{t-1} + B_t 其实也可以放进这个框架。我们把每个位置t对应一个二元组(A_t, B_t),它描述的是“从t-1时刻状态到t时刻状态的一次迁移”。如果我有两个相邻区间,左侧区间整体迁移是(M_L, N_L),右侧区间整体迁移是(M_R, N_R),那么把两个区间拼起来后的整体迁移是:
M = M_R M_L N = M_R N_L + N_R
这个组合规则就是整个并行扫描的基石。只要这个组合可以高效计算,我们就能把L个位置的递推,通过log L轮的树形归约并行化,这也就是所谓的“并行扫描”。在这里,M是D×D矩阵,N是D×D矩阵,组合一次的开销是两次矩阵乘加一次矩阵加。
对GDN来说,每个位置的A_t是“对角+秩一修正”,位置不同、修正方向也不同。连续乘一堆这样的矩阵,结果通常是一个D×D稠密矩阵,没有简单的低秩闭合形式,所以组合算子必须真刀真枪地做矩阵乘法。
2.2 Chunk-wise:把绝大多数工作换成矩阵乘
纯并行扫描虽然漂亮,但实际工程里很少有人直接对L个位置做全局scan,原因后面会细讲。更主流的选择是先分块(chunk),再用矩阵乘法消化块内的计算。
思路和FlashAttention的分块思想类似:把序列切成若干长度C的chunk。对每个chunk,我们事先算出这个chunk作为一个整体的等效迁移(M_b, N_b),然后chunk与chunk之间只传递这两个D×D矩阵,而不是把chunk内部每个位置的中间状态都暴露出来。
块间关系是:
S_{b+1} = M_b S_b + N_b
其中b是chunk序号。如果序列长度L=32768,chunk大小C=64,那么chunk数量只有512。这个规模下的块间递推,无论是串行循环还是再一次扫描,开销都小得多。绝大多数计算量发生在chunk内部,而chunk内部因为长度固定且较短(通常64或者128),可以用共享内存和矩阵乘来加速。
2.3 两条路线怎么选
我实际做的时候体会很深:纯并行扫描的问题是“work效率”差。它虽然把时间复杂度压到了O(log L)轮,但总的计算量是O(L log L)级别,比原来的O(L)递推还多。而且扫描的每一轮都要交换共享内存里的数据,如果状态矩阵是D×D,D又取64或者128,那一轮交换的数据量非常可观。
chunk方法则不同:块内的主要计算可以通过矩阵乘(TensorCore)来完成,块间只处理很少的块数量。它本质上是“用更多并行但更少冗余”的方式处理依赖。实际实现中,chunk大小一般选64或128。太小,块间数量太多;太大,块内共享内存和计算压力吃不消。
3. 一步一步实现:从Python参考到CUDA/Triton
3.1 先写一个能跑的Python参考
不管你要写多快的kernel,第一步永远是先写一个绝对正确的朴素实现,用来对拍。我一般先写一个纯循环版,形状按(D, D)状态、行向量token来约定。为了描述方便,假设q、k、v的形状都是(seq_len, dim),状态S是(dim, dim)矩阵,输出也是(seq_len, dim)。
import torch def gdn_naive(q, k, v, alpha, beta): # q, k, v: (L, D); alpha, beta: (L,) 范围(0,1) L, D = q.shape S = torch.zeros(D, D, dtype=q.dtype, device=q.device) y = torch.empty_like(q) for t in range(L): y[t] = q[t] @ S A = torch.eye(D, device=q.device) \ - alpha[t] * beta[t] * torch.outer(k[t], k[t]) B = alpha[t] * beta[t] * torch.outer(k[t], v[t]) S = A @ S + B return y这个版本笨是笨,但它把递推讲清楚了。建议你先拿极小维度(比如L=8, D=8)跑通,确认输出与标准实现一致,再往下走。
3.2 Chunked版的参考实现
接下来是chunked版本。这个版本虽然还是Python循环,但它已经把“块内计算”和“块间状态迁移”分开了,结构上跟GPU kernel是同一个骨架。
def gdn_chunked(q, k, v, alpha, beta, C=64): L, D = q.shape S = torch.zeros(D, D, dtype=q.dtype, device=q.device) y = torch.empty_like(q) n_chunks = (L + C - 1) // C for b in range(n_chunks): start = b * C end = min(start + C, L) qb, kb, vb = q[start:end], k[start:end], v[start:end] ab, bb_ = alpha[start:end], beta[start:end] # 块内需要保存每个位置的 partial (M, N) M_part = [torch.eye(D, device=q.device)] # M_part[0] = I N_part = [torch.zeros(D, D, device=q.device)] for i in range(end - start): A = torch.eye(D, device=q.device) \ - ab[i] * bb_[i] * torch.outer(kb[i], kb[i]) B = ab[i] * bb_[i] * torch.outer(kb[i], vb[i]) # 累积:代表从块起始位置的状态出发,作用到当前位置的迁移 M_part.append(A @ M_part[-1]) N_part.append(A @ N_part[-1] + B) for i in range(end - start): y[start + i] = qb[i] @ (M_part[i + 1] @ S + N_part[i + 1]) # 块间更新全局状态 S = M_part[-1] @ S + N_part[-1] return y这个版本里,M_part和N_part就是在内存中显式地做“块内并行扫描”的结果。真正的CUDA kernel不会像这样在Python里存一堆D×D矩阵,但逻辑是完全一样的:你要在所有位置上算出从块起始点到当前位置的累积迁移,然后再用它来组合全局状态。
3.3 关键步骤怎么转化为并行
现在到了最关键的部分:怎么把上面这个循环重现成一个高性能GPU实现。
第一步,块内partial (M, N)的计算。长度为C的块内,其实就是在做一次长度为C的并行扫描。理想情况下,我们希望在log C轮内通过扫描树算出所有位置的M_part和N_part。可问题在于每个partial都是D×D矩阵,如果D=64、C=64,就是65个64×64矩阵。放FP32大约是646465*4 ≈ 1MB,超过了常见的shared memory限制。所以工程上不会简单地把整个partial矩阵塞进shared。
常用的处理方式并不是“不做扫描”,而是压缩要扫描的对象。注意最终使用的其实只有两个东西:一是每个位置i的累积极迁移对初始状态S_b的读取效果,即 q_i^T M_part[i+1];二是每个位置i的累积新增贡献,即 q_i^T N_part[i+1]。这两个都是D维度向量,而不是D×D矩阵。所以可以在扫描过程里只维护与q_i相关的向量版本,但这需要调整扫描顺序,并不是每一组q_i都能直接由前一个位置的partial推出来。这也是为什么很多现成kernel采用了一种更工程化的写法:先通过矩阵乘计算所有“来自S_b的读取”,再把块内历史贡献用一个上三角掩码的矩阵乘来算。
这里我给一个更接地气的视角。假设块内C=64,张量形状都是[C, D]。我们可以一次性计算所有位置对初始状态的投影:
O_init = Q_b @ S_b # [C, D]
这个用一次矩阵乘就能解决。它代表的是“如果状态不随时间变化,每个位置能读到什么”。剩下的问题是块内产生的增量如何传播到后面的位置。这部分可以看作一个“三角形”的权重矩阵:位置j的写入,经过中间位置的A衰减,最终被位置i读取,i >= j。GDN里这个传递权重不是简单的标量门控乘积,因为它涉及到矩阵A的连乘。实际实现时,块内C=64这个规模下,可以用共享内存配合warp内shuffle做一轮小范围的扫描,或者直接把C×D的中间投影在寄存器里滚动更新。
不过说实话,手写这部分非常容易写到逻辑混乱。我建议的做法是:先用共享内存维护块内partial的“作用向量”,也就是上面说的q_i^T M_part和q_i^T N_part,用C个线程、每个线程负责一个位置,配合线程束内扫描逐步扩散结果。这样虽然不如TensorCore快,但至少已经避开了D×D矩阵在块内爆炸的存储问题,可以先把正确性和性能轮廓跑出来。
第二步,块间的状态传播。块数通常只有L/C那么多,比如L=8192、C=64,块数就是128。这种量级下,块间递推可以直接串行,因为每次更新只需要一次D×D的矩阵乘法S = M_b @ S,128次矩阵乘相对整个kernel来说开销不大。如果想进一步并行,也可以把块的(M_b, N_b)序列再做一次并行扫描,但收益有限,还多了延迟,我一般不建议在块数少于几百时强行并行。
第三步,kernel融合。完整的计算流程是:加载QKV和门控,算块内partial,算块间状态,算输出,写回。这里最影响性能的是中间结果是否被打到全局内存。尽量把“算partial”和“算输出”放在同一个kernel里。块间状态传播如果串行,可以放在kernel末尾由最后一个block或一个小kernel完成。更常见的方案是拆成两个kernel:第一个kernel负责每个chunk的partial和等效迁移,第二个kernel负责按块顺序更新全局状态并计算输出。第一个kernel并行度极高,第二个kernel虽然串行但计算量小。
3.4 参数与精度要点的具体建议
chunk大小C的选择。我试过32、64、128几个值。C=32时块间数量偏多,块内扫描树深度小,但矩阵乘的形状太小,TensorCore利用率低;C=128时块内shared memory压力大,而且块内partial的warp级扫描链变长,容易成为瓶颈。综合下来C=64是大多实现里比较稳的选择。如果你的状态维度D比较小(比如32以下),C可以适当放大到128;D大(比如128)时,C=32甚至更小反而更好。
状态矩阵S建议始终用FP32维护。GDN的更新里有一个明显的减法效应:k_t^T S_{t-1}会把旧状态中的一部分“擦掉”。在FP16下做这样的减法,累积误差会很讨厌,尤其长序列。我的做法是:Q、K、V和门控以FP16/BF16存储和读取,但在状态更新路径上,也就是A@S+B这一步,用FP32计算。实测下来,这比全FP32省显存,数值稳定性也够用。
精度验证时,用torch.testing.assert_close对比朴素递推和chunked实现,rtol建议1e-4、atol建议1e-5起步。如果差太多,先查维度方向是否反了,再查门控的alpha和beta有没有被sigmoid限制。
4. 扩展到Mamba、RWKV、GLA等Linear Attention
4.1 Mamba的scan与chunk化
Mamba的状态更新比GDN简单不少:A是固定的对角矩阵,B_t和u_t都是向量。因为A是对角的,状态S的每个维度可以独立更新,矩阵乘退化为逐元素乘加。这导致Mamba非常容易用并行扫描实现,也是Mamba官方kernel速度快的原因之一。
chunk化Mamba时,块内部的partial (M,N)会简单得多:M是C个标量衰减的累积,N是带衰减的累积输入。你甚至可以把一条chunk内的所有位置组合成一个上三角衰减矩阵,然后用一次矩阵乘算出整个chunk的输出。这也是Mamba2论文里展示的chunked selective scan思想。
4.2 RWKV与GLA的不同并行方式
RWKV的时间混合本质上是一个带指数衰减的线性attention权重。它的并行实现可以走线性attention的路线:把所有位置的状态压缩求和,再对query逐个读取。不过RWKV的衰减系数是全局常数,做并行特别友好。你可以直接构造一个衰减矩阵,把所有历史位置对当前query的贡献用矩阵乘算出来。
GLA(Gated Linear Attention)加了一个逐维度的门控向量,所以衰减不是标量,而是向量。它的chunk化跟GDN很像,只不过A_t是对角矩阵,不是“单位阵加秩一修正”。对角矩阵的好处是每个维度独立,partial的扫描复杂度低一些。GLA的官方实现里就用了chunked形式,块内通过矩阵乘计算带门控衰减的KV累积,块间再做一次轻量扫描。
4.3 一张表看清各架构该怎么并行
| 模型 | 状态更新 | 衰减形式 | GPU并行主流方案 |
|---|---|---|---|
| Mamba | S_t = A S_{t-1} + B_t u_t | 对角矩阵,固定 | 并行扫描 + chunked矩阵乘 |
| RWKV | 指数衰减线性attention | 标量常数 | 线性核 + 衰减矩阵乘 |
| GLA | S_t = diag(α_t) S_{t-1} + k_t v_t^T | 对角门控,逐位置变化 | chunked矩阵乘 + 块间扫描 |
| DeltaNet | S_t = (I - β k k^T) S_{t-1} + β k v^T | 单位阵+秩一修正 | chunked扫描,块内矩阵乘 |
| GDN | 门控+delta rule | 门控缩放后的秩一修正 | 同上,注意门控与擦除的融合 |
这张表我建议贴在工位旁。做并行实现前,先看你的模型属于哪一类,再决定scan的算子长什么样。
5. 常见问题、排查与性能调优实录
5.1 为什么我自己写的scan比FlashAttention还慢
这是最容易劝退新手的现象。明明线性注意力复杂度更低,为什么自己按论文写的并行扫描比FlashAttention还慢?
原因几乎总是“work效率太低+访存太大”。朴素的并行扫描,每一轮都要把整个状态矩阵在共享内存和寄存器之间搬一次,D=64时,状态是4096个浮点数,一个block里每轮都要处理这么多数据,而FlashAttention那种分块矩阵乘主要靠TensorCore,访存模式也更规则。
改进方向就是前面说的chunk。块内不要用全局扫描,先把能并行投影的部分用矩阵乘算掉,剩下依赖链很短的小段扫描。当你看到L=8192、D=64时速度终于超过标准Attention,你会明白:Linear Attention的GPU实现拼的不是“能不能并行”,而是“怎么把递归压成长度很短的一小段”。
5.2 精度问题:delta rule在FP16下的坑
GDN是Linear Attention里对精度最不友好的一种,因为它的Delta Rule本质上在做一个减法:你告诉状态“把和这个key方向重合的一部分记忆擦掉”。在FP16下,状态S的数值范围如果比较大,k_t^T S_{t-1}的结果会被舍入,再拿这个结果去乘k_t,误差会被放大。我踩过的坑是:序列长度只有4096,D=64,FP16跑出来的最终状态和FP32参考版本相差0.5,输出层完全崩了。
解决方法是分层精度:Q、K、V可以低精度,状态路径必须高精度。如果你的kernel里全部用了FP16,先改成状态矩阵FP32,对比一下数值误差立刻小两个数量级。想在FP16下做更极致的优化,可以每个chunk结束后对状态做一个重缩放(rescale),把天量级状态压回到合理范围,但这样要额外记录缩放因子,复杂度上升,我的建议是不到万不得已不需要。
5.3 显存与训练时的额外开销
GDN推理时只需要维护一个D×D状态矩阵,看着不大。但训练时反向传播要存下每个位置的中间状态或者partial,不然没法反算梯度。如果L=32768、D=64、batch=8、head=8,要存的partial就是非常大的开销,甚至比FlashAttention的注意力矩阵还占显存。
我的经验是训练时用梯度检查点(gradient checkpointing),只在chunk边界存状态,块内partial用重计算。这样显存占用大概降到原来的四分之一到五分之一,代价是多一次前向重算,但显存不足导致的训练中断更伤。
5.4 调试手段与基准测试方法
我调试这类kernel的固定流程:先用朴素Python实现跑小shape,再写chunked版本做一致性对比;确认无误后,把chunked版本里的块内循环改成一个独立的scan函数,逐步替代成扫描树;最后做kernel融合和内存优化。
性能测试不要只看总时长。要拆开看:如果矩阵乘占比很低,时间都花在访存和扫描同步上,说明chunk切得不好或者块内矩阵乘形状太小。用Nsight Compute看shared memory bank conflict和occupancy。注意Linear Attention的算子很容易被shared memory限制并发,不是计算瓶颈。
5.5 问题速查表
| 现象 | 可能原因 | 解法 |
|---|---|---|
| 数值与参考实现偏差大 | 状态用了FP16 | 状态路径改FP32 |
| 长序列速度不如预期 | chunk内扫描太重 | 增大chunk内矩阵乘占比 |
| 显存爆炸 | 中间partial存太多 | 梯度检查点 |
| block内shared memory不够 | C过大或D过大 | C降到32/64,或压缩M/N |
| 输出后半段全错 | 块间状态传播顺序写错 | 单独验证S_b的更新 |
| 性能在batch小时很差 | 并行规模被chunk数限制 | 多个chunk合并到一个block处理 |
6. 最后的一点实践体会
把我的经验浓缩成两句:第一,先接受“串行依赖没法完全消失”,GPU上能做的是把依赖锁进很小的chunk,让矩阵乘成为主角;第二,数值稳定性永远比性能优先,Linear Attention跑出了惊人的数字但状态发散,这模型没法用。
我自己写这类算子时最吃亏的地方,是总想追求一个“从公式到极致kernel”的完美方案,结果在scan树和矩阵乘之间反复横跳。后来学乖了:先写一个能跑能验证的chunked版本,不管它多慢,再针对profile结果一步步优化。另外一个小技巧——调C时不要只看快慢,把C=64和C=128的“块内partial精度”也对比一下,因为C越大,块内的累积路径越长,对beta和alpha这种小数值越敏感,C=128在FP16下有时数值误差真的会明显变大。
如果你正在把GDN或者类似的Linear Attention往GPU上搬,希望这篇文章能帮你少走点弯路。沿着“状态递推结构化 -> chunk化 -> 块内矩阵乘 -> 块间scan/串行”这条路走,就算手写不出最顶级的kernel,至少能有一个稳定可用的高性能版本。