FlashQLA流水线解析

鱿鱼圈 Lv4

链接: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 在解决什么问题

如果把它写成最直白的版本,逻辑大概会是:

  1. 先读入一个 chunk 的 Q/K/V/A/g/b
  2. 用旧状态 S_old 算一些中间量
  3. 算出当前 chunk 的输出 O_chunk
  4. 算出新的状态 S_new
  5. O_chunkS_new 写回
  6. 再处理下一个 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
2
3
4
5
6
q_shared = T.alloc_shared((2, block_S, DK), ...)
k_shared = T.alloc_shared((2, block_S, DK), ...)
v_shared = T.alloc_shared((2, block_S, block_DV), ...)
a_shared = T.alloc_shared((2, block_S, block_S), ...)
g_shared = T.alloc_shared((2, block_S), ...)
b_shared = T.alloc_shared((2, block_S), ...)

这说明 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_0
  • bar_1
  • bar_3
  • bar_4
  • bar_5
  • bar_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
2
3
T.barrier_wait(data_is_ready[i_s % 2], ...)
T.barrier_arrive(bar_0)
T.barrier_wait(bar_0, i_s % 2)

意思是:

  1. 先等 Producer 说“这个 shared slot 的数据装好了”
  2. 再让 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
2
g_exp_shared[j] = exp2(g[j] * 1.442695)
g_rev_exp_shared[j] = exp2((g_last - g[j]) * 1.442695)

这一步是把 gate 相关的指数形式提前算好。

O-path 做什么

1
2
3
P = QK^T
G = lower-triangular gate matrix
Ag = G * A * b

也就是提前准备当前 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 条逻辑链揉在一起:

  1. 状态链S <- gate * S + ...
  2. 值修正链U -> W -> Vd -> V'
  3. 输出链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 1 ping-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_1h_sharedg_exp_shared 等已经准备好
  • bar_4vd_shared 准备好了,O-path 可以拿去算 Pg @ Vd
  • bar_o:所有 chunk 算完后,最后一个 o_shared 可以安全写回

17.2 为什么有的路径先 T.copy(..., shared) 再 GEMM

因为 CTA 内不同线程组要共享这些中间量。

如果某个量只在寄存器里,别的线程组看不到。
放到 shared 后,另一路线程才能继续消费。


17.3 为什么不是直接一个线程组把所有事都做完

因为这样:

  • 读数据时算力浪费
  • 算输出时 store 线程空着
  • 整体重叠度差

现在这种拆法,等于是用编排复杂度换吞吐。


18. 用一句人话概括整个流水线

如果完全不用公式,而用人话描述:

这个 kernel 把一个大任务拆成四个工种: 一组人专门维护状态, 一组人专门准备状态更新量, 一组人专门计算输出, 另一组人专门负责搬下一批原料和把上一批结果送出去。
大家围绕同一个 chunk 分阶段协作,同时还用双缓冲让“下一块数据的加载”和“当前块计算”重叠起来。

这就是它的多级流水机制。


19. 一个可以记住的总图

你可以把每个 chunk 的一轮计算记成下面这张抽象图:

1
2
3
4
5
6
7
8
9
Producer:    load(Q/K)  load(V/b)  load(A/g) ---------------------->

S-path: read S -> decay S -------------> S += K^T @ V'

V-path: precompute g -> U -> W -> Vd -> V'

O-path: P/G/Ag -> Q@S -> Pg -> O += Pg@Vd -> O_shared

Store: <------------------- store old O / old H ----------------

再叠加双缓冲后,时间轴上就会变成:

1
2
3
chunk i-1 : store
chunk i : compute
chunk i+1 : preload

虽然不是完全理想重叠,但设计目标就是尽量往这个方向靠。


20. 最后一句总结

tilelang_fused_chunk_gdr_fwd_kernel 的多级流水机制,本质就是:

把一个 chunk 的 forward 拆成“状态更新、值修正、输出生成、数据搬运”四条协作链,再通过双缓冲和多级 barrier,让这些链在一个 CTA 内尽量重叠执行,从而减少空转、提高片上复用和整体吞吐。

如果你后面还想继续深挖,最推荐的下一步不是继续看所有 barrier,而是:

  1. 先把 U / W / Vd / V' / O / S 的数学关系真正看懂
  2. 再回头看 barrier,就会觉得它们只是“把这些公式安排到不同线程组上执行”的调度工具

附:阅读这份 kernel 的建议顺序

如果你之后还要继续自己读源码,建议按下面顺序:

  1. 先看 CTA 分工:tx < 128 / 256 / 384 / else
  2. 再看 shared / fragment 变量分别对应什么数学对象
  3. 再看一个 chunk 内 i_s 循环里三条路径分别做了什么
  4. 最后再看 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 进行许可。
评论
目录
FlashQLA流水线解析