FlashQLA流水线解析
链接:Qwen
tilelang_fused_chunk_gdr_fwd_kernel 流水线机制详解
本文面向第一次接触这个 kernel 的读者,目标不是逐行翻译代码,而是帮助你建立一个清晰的整体心智模型:
- 这个 kernel 在算什么
- 为什么它要拆成多组线程
- 为什么要做双缓冲
- 为什么会有这么多 barrier
- 一个 chunk 在 kernel 里到底经历了哪些阶段
文中默认讨论的是:
flash_qla/ops/gated_delta_rule/chunk/hopper/fused_fwd.py- kernel 名称:
tilelang_fused_chunk_gdr_fwd_kernel - 不重点展开 CP 分支,只把它当作 shape / 写回控制的一部分
1. 先用一句话理解这个 kernel
这个 kernel 的本质是:
对某个
(序列段, head, value-dim 分块),沿着时间按 chunk 逐块扫描;在扫描过程中,一边维护递推状态S,一边计算当前 chunk 的输出O,并且尽量把中间量留在片上(shared / fragment)立即消费,而不是写回 global memory。
这就是它和“多 kernel 拆开做”的最大不同。
2. 这个 kernel 在解决什么问题
如果把它写成最直白的版本,逻辑大概会是:
- 先读入一个 chunk 的
Q/K/V/A/g/b - 用旧状态
S_old算一些中间量 - 算出当前 chunk 的输出
O_chunk - 算出新的状态
S_new - 把
O_chunk、S_new写回 - 再处理下一个 chunk
这种“串行版”虽然直观,但会有几个问题:
- 读数据的时候,计算线程在等
- 算状态的时候,输出线程可能在等
- 写回的时候,大量线程又在等
- chunk 与 chunk 之间几乎没有重叠
所以这个 kernel 的目标是:
让“加载下一块数据”“计算当前块”“写回上一块结果”尽量同时发生。
这就是流水线设计的出发点。
3. CTA 的基本分工
kernel launch 是:
1 | with T.Kernel(T.ceildiv(DV, block_DV) * batch_size * H, threads=512) as (bbhv,): |
意思是:
- 一个 CTA 负责一个
batch/head/value-block DV维如果太大,会按block_DV再切成多个子块
所以一个 CTA 关注的是:
- 哪条序列 / 哪个 varlen 子段
- 哪个 head
- 哪一段
DV
然后这个 CTA 再沿时间维 chunk by chunk 处理。
4. 一个 CTA 里的 512 个线程怎么分工
这个 kernel 不是 512 个线程都做同一件事,而是硬切成四组:
4.1 线程 0 ~ 127:状态路径(S-path)
负责:
- 维护 recurrent state
S - 从
S_old生成S_new - 最后写
final_state
核心公式可以粗略理解成:
1 | S <- g_last * S + K^T @ V' |
4.2 线程 128 ~ 255:值修正路径(V-path)
负责:
- 先从当前状态算出
U = K @ S - 再构造
W = V - g * U - 再算
Vd = Ag @ W - 再算
V' = (g_last / g) * Vd
这一组线程的工作,是给状态更新路径准备“状态增量”。
4.3 线程 256 ~ 383:输出路径(O-path)
负责:
- 算
P = QK^T - 构造门控下三角矩阵
G - 构造
Ag = G * A * b - 算
O = Q @ S - 再加上 chunk 内修正项
Pg @ Vd
所以它负责真正的 token 输出。
4.4 线程 384 ~ 511:Producer / Store 路径
这一组线程不参与主要数学计算,而是负责:
- 读取下一个 chunk 的
Q/K - 读取下一个 chunk 的
V/beta - 读取下一个 chunk 的
A/g - 写回上一块的
O - 写回上一块的中间状态
h
你可以把它看作 CTA 内专门负责“搬数据”和“落结果”的后勤线程组。
5. 共享内存为什么很多 buffer 都是 (2, ...)
你会看到很多 shared buffer 都带第一维 2:
1 | q_shared = T.alloc_shared((2, block_S, DK), ...) |
这说明 kernel 使用了 双缓冲(double buffering / ping-pong buffering)。
直观理解
slot 0:当前 chunk 正在计算slot 1:下一 chunk 正在加载- 下一轮交换
索引通常是:
1 | i_s % 2 |
也就是:
- 偶数 chunk 用 buffer 0
- 奇数 chunk 用 buffer 1
这样 Producer 可以在 Consumer 计算 chunk i 的同时,把 chunk i+1 的数据预取到另一个 slot。
6. 为什么要有那么多 barrier
这是小白最容易被劝退的部分,但其实可以按功能分成两类看。
6.1 第一类:双缓冲读写协调
data_is_ready
表示:
某个 buffer slot 的输入数据已经装好了,计算线程可以开始用。
data_is_free
表示:
某个 buffer slot 的数据已经被消费完,Producer 可以安全重用这块 buffer。
这两个 barrier 是双缓冲的基础。
6.2 第二类:chunk 内部不同线程组之间的依赖同步
例如:
- O-path 要等 V-path 先算好
Vd - V-path 要等 S-path 先把
S搬到 shared - Store 线程要等 O-path 真正把
o_shared填好
这些同步点就用到了:
bar_0bar_1bar_3bar_4bar_5bar_o
你可以先不用记每个 barrier 的名字,而是先记一句话:
这些 barrier 的目的,就是让 CTA 内不同角色线程组在“数据刚好准备好”的那个时刻接上,而不是过早读到脏数据。
7. 一个 chunk 在 kernel 里经历了哪些阶段
虽然代码里的注释都写成 STAGE 0,但实际上每个 chunk 内部被拆成了多个小阶段。为了理解方便,我们把它们重新编号成:
- Stage A:本 chunk 输入 ready
- Stage B:准备状态与 gate 派生量
- Stage C:生成
U / W / P / G - Stage D:生成
Vd / Pg - Stage E:生成
V'、累加输出 - Stage F:更新
S,提交O
下面分别说。
8. Stage A:本 chunk 输入 ready
所有 consumer 线程组一开始都会做:
1 | T.barrier_wait(data_is_ready[i_s % 2], ...) |
意思是:
- 先等 Producer 说“这个 shared slot 的数据装好了”
- 再让 S/V/O 三组消费者统一从这个 chunk 的入口开始
你可以把这里理解成:
chunk
i的“发车信号”
9. Stage B:准备状态和 gate 派生量
S-path 做什么
1 | T.copy(h_fragment, h_shared) |
把当前状态 S 从 fragment 拷到 h_shared,让别的线程组也能读。
V-path 做什么
1 | g_exp_shared[j] = exp2(g[j] * 1.442695) |
这一步是把 gate 相关的指数形式提前算好。
O-path 做什么
1 | P = QK^T |
也就是提前准备当前 chunk 输出所需的局部矩阵。
10. Stage C:各条计算主线开始展开
V-path:先算 U = K @ S
1 | U = K @ S |
然后:
1 | W = V - g * U |
这一步的直觉是:
- 当前状态
S可以先预测一部分V - 然后用
V - g*U得到需要被“新信息”补进去的残差部分
O-path:先算状态项
1 | O = Q @ S |
这相当于先算出“旧状态对当前输出的贡献”。
S-path:先衰减旧状态
1 | S = g_last * S |
表示状态会先经历一个衰减,再加上本 chunk 的新贡献。
11. Stage D:构造 chunk 内修正量
V-path:算 Vd = Ag @ W
1 | Vd = Ag @ W |
这是当前 chunk 局部解的核心。
O-path:算 Pg = scale * G * P
1 | Pg = scale * G * P |
这一步是把 chunk 内 token 间交互矩阵 P 乘上 gate 结构和 scale,变成输出修正项的左侧矩阵。
12. Stage E:生成状态更新增量和输出修正量
V-path:算 V'
1 | V' = (g_last / g) * Vd |
这一步的结果 V' 是给状态递推路径用的。
O-path:累加 chunk 内输出修正
1 | O += Pg @ Vd |
所以当前输出不是只有 Q @ S,而是:
1 | O = Q @ S + Pg @ Vd |
也就是:
- 一部分来自旧状态
- 一部分来自当前 chunk 局部交互的增量修正
13. Stage F:更新最终状态并提交结果
S-path:最终更新状态
1 | S += K^T @ V' |
所以状态更新可以写成:
1 | S_new = g_last * S_old + K^T @ V' |
O-path:把当前 chunk 的输出放到 o_shared
1 | T.copy(o_fragment, o_shared) |
之后 store 线程组会负责把它真正写回 global memory。
14. 为什么这套流水看起来很“碎”
因为它不是一个单一 GEMM kernel,而是把 3 条逻辑链揉在一起:
- 状态链:
S <- gate * S + ... - 值修正链:
U -> W -> Vd -> V' - 输出链:
P -> G -> Pg -> O
如果把它们完全串行写,很多线程会长期闲着。
所以作者把它拆开后,让三组线程:
- 在同一个 chunk 内做不同子任务
- 通过 barrier 在必要的时候同步
- 中间量尽量放 shared/fragment 上直接消费
这就形成了所谓的“多级流水”。
15. Producer / Store 线程组到底在重叠什么
最后那组线程(tx >= 384)非常关键,它让 kernel 不只是“chunk 内流水”,而是有了 chunk 间重叠。
它做的事包括:
15.1 预取下一个 chunk 的输入
分成三小组:
- 32 线程加载
Q/K - 32 线程加载
V/beta - 32 线程加载
A/g
所以当 S/V/O 三组线程在算 chunk i 时,Producer 已经在往另一个 shared slot 里装 chunk i+1。
15.2 写回上一个 chunk 的结果
还有一组线程负责:
- 把上一 chunk 的
o_shared写回o - 把上一 chunk 的
h_shared写回h
这意味着从全局看,CTA 内部近似在做:
- load chunk i+1
- compute chunk i
- store chunk i-1
三件事情部分重叠。
16. 这就是“多级流水”的真正含义
所以这个 kernel 的流水不是一层,而是至少有三层:
第一层:chunk 级双缓冲流水
slot 0/slot 1ping-pong- 让下一 chunk 的数据预取和当前 chunk 的计算重叠
第二层:chunk 内多阶段流水
- 一个 chunk 内被拆成多个子阶段(B/C/D/E/F)
- 各阶段之间通过 barrier 连接
第三层:CTA 内角色流水
- S-path
- V-path
- O-path
- Producer/Store-path
四组线程不是同时干同一件事,而是像四条小流水线一样协作。
17. 小白最容易迷糊的地方
17.1 为什么 bar_0 / bar_1 / bar_3 / bar_4 / bar_5 / bar_o 这么多
不要把它们当成“很多无意义同步”。
把它们理解成:
- 某个共享变量准备好了
- 需要通知另一组线程来取
例如:
bar_1:h_shared、g_exp_shared等已经准备好bar_4:vd_shared准备好了,O-path 可以拿去算Pg @ Vdbar_o:所有 chunk 算完后,最后一个o_shared可以安全写回
17.2 为什么有的路径先 T.copy(..., shared) 再 GEMM
因为 CTA 内不同线程组要共享这些中间量。
如果某个量只在寄存器里,别的线程组看不到。
放到 shared 后,另一路线程才能继续消费。
17.3 为什么不是直接一个线程组把所有事都做完
因为这样:
- 读数据时算力浪费
- 算输出时 store 线程空着
- 整体重叠度差
现在这种拆法,等于是用编排复杂度换吞吐。
18. 用一句人话概括整个流水线
如果完全不用公式,而用人话描述:
这个 kernel 把一个大任务拆成四个工种: 一组人专门维护状态, 一组人专门准备状态更新量, 一组人专门计算输出, 另一组人专门负责搬下一批原料和把上一批结果送出去。
大家围绕同一个 chunk 分阶段协作,同时还用双缓冲让“下一块数据的加载”和“当前块计算”重叠起来。
这就是它的多级流水机制。
19. 一个可以记住的总图
你可以把每个 chunk 的一轮计算记成下面这张抽象图:
1 | Producer: load(Q/K) load(V/b) load(A/g) ----------------------> |
再叠加双缓冲后,时间轴上就会变成:
1 | chunk i-1 : store |
虽然不是完全理想重叠,但设计目标就是尽量往这个方向靠。
20. 最后一句总结
tilelang_fused_chunk_gdr_fwd_kernel 的多级流水机制,本质就是:
把一个 chunk 的 forward 拆成“状态更新、值修正、输出生成、数据搬运”四条协作链,再通过双缓冲和多级 barrier,让这些链在一个 CTA 内尽量重叠执行,从而减少空转、提高片上复用和整体吞吐。
如果你后面还想继续深挖,最推荐的下一步不是继续看所有 barrier,而是:
- 先把
U / W / Vd / V' / O / S的数学关系真正看懂 - 再回头看 barrier,就会觉得它们只是“把这些公式安排到不同线程组上执行”的调度工具
附:阅读这份 kernel 的建议顺序
如果你之后还要继续自己读源码,建议按下面顺序:
- 先看 CTA 分工:
tx < 128 / 256 / 384 / else - 再看 shared / fragment 变量分别对应什么数学对象
- 再看一个 chunk 内
i_s循环里三条路径分别做了什么 - 最后再看 barrier —— 不要一开始就盯着 barrier 看
这样更容易读懂。
- 标题: FlashQLA流水线解析
- 作者: 鱿鱼圈
- 创建于 : 2026-06-24 23:50:00
- 更新于 : 2026-06-30 22:51:32
- 链接: https://yuyanqi.com/2026/06/24/FlashQLA流水线解析/
- 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。