KDA源码剖析之FlashKDA(下)
前置知识:KDA算法公式、triton、cuda、c++
仓库链接:MoonshotAI/FlashKDA: FlashKDA: high-performance Kimi Delta Attention kernels
公式回顾
kernel 2 代码
1 |
|
计算流程图


本文目标:从底层 MMA 指令出发,解释 Phase 1-6 中所有寄存器操作的"为什么"。
1 MMA 寄存器基础知识
1.1 MMA 指令基础
GPU 的矩阵乘法不是一个线程算一个元素,而是 一个 warp (32线程) 协作算一整个矩阵块。
SM80+ 的核心指令:mma.sync.aligned.m16n8k16.f32.bf16.bf16.f32
含义:
1 | C[16×8] += A[16×16] × B[16×16]^T |
注意:
- 一条 mma 指令只能算 16×8 的输出
- 要算 16×16 的输出,需要 两条 mma 指令(列方向拼接:8+8=16)
- CuTe 的一次
gemm()调用 = 两条 mma 指令 = 得到 16×16 的 C
1.2 三种格式:A格式、B格式、C格式
MMA 指令要求 A、B、C 三个操作数在 32 个线程的寄存器中以特定方式分布。三种分布方式完全不同!这是理解所有寄存器操作的关键。
1.2.1 C 格式(累加器格式 / 输出格式)
C 是 MMA 的输出。16×16 = 256 个元素,分给 32 线程,每线程 8 个。
每线程持有的 8 个 fp32 元素在矩阵中的位置:
1 | groupID = lane_id / 4 (0..7) |
画出来(以 groupID=1, tid_in_grp=1 即 lane_id=5 为例):
1 | 16×16 矩阵: |
特点:
- 每线程跨两行(row 和 row+8)
- 每线程在一行中持有连续两列
- 同行同列的 4 个线程(tid_in_grp=0,1,2,3)覆盖了该行的全部 8 列
这就是为什么:
beta0对应 row_half=0 的 4 个元素(groupID 那行)beta1对应 row_half=1 的 4 个元素(groupID+8 那行)- 每行的 beta 值相同,一个线程只碰两行,所以只需要 2 个标量
1.2.2 A 格式(左操作数格式)
A 是 16×16 的左矩阵。每线程也持有 8 个 bf16(打包成 4 个 uint32)。
- 行:
lane_id % 16(0…15,每线程只在一行) - 列:分成4组,由
lane_id / 16和具体分配决定
关键区别:A 格式中每线程持有同一行的元素(按 K 维度排列),C 格式中每线程持有两行的元素(按输出行排列)。
1.2.3 B 格式(右操作数格式)
B 是 16×16 的右矩阵(实际是转置后使用)。每线程也持有 8 个 bf16,但排列方式又不同于 A 和 C。
每线程持有同一列的元素(按 K 维度排列)。
1.2.4 三种格式的核心区别
| 格式 | 每线程持有的几何含义 | 典型操作 |
|---|---|---|
| A 格式 | 同一行的若干元素 | LDSM_N 加载,作为 gemm 左操作数 |
| B 格式 | 同一列的若干元素 | LDSM_N 加载,作为 gemm 右操作数 |
| C 格式 | 两行的若干元素 | gemm 输出,逐元素运算,STSM 存储 |
三种格式是不可互换的!如果你在 C 格式的寄存器中有数据,想把它作为 B 操作数喂给下一个 gemm,你必须做格式转换(C→B)。这就是 MOVM_T 存在的意义。
1.3 从 smem 到寄存器:LDSM 指令
LDSM = Load Shared Memory(ldmatrix PTX 指令)
一条 LDSM 指令:32 个线程协作,从 shared memory 读取一个矩阵块,自动把数据按 A/B/C 格式分配到各线程的寄存器中。
1.3.1 LDSM_N(Normal)
按正常方向加载。
copy_A + LDSM_N:把 smem 中 16×16 块加载为 A 格式copy_B + LDSM_N:把 smem 中 16×16 块加载为 B 格式copy_C + LDSM_N:把 smem 中 16×16 块加载为 C 格式
1.3.2 LDSM_T(Transposed)
加载时做转置。
copy_A + LDSM_T:把 smem 中 16×16 块转置后加载为 A 格式copy_C + LDSM_T:把 smem 中 16×16 块转置后加载为 C 格式
物理数据不变,但加载后在寄存器中的行列关系变了。
应用:Phase 6 中 s_acc 的物理布局是"列优先"(为 Phase1 的 B 格式设计),但 Phase 6 需要按行读取,所以用 LDSM_T 转置加载。
1.3.3 对应 Phase 中的用法
| Phase | LDSM 操作 |
|---|---|
| Phase 1 | LDSM_N 加载 kd, qd → A 格式(tCrA_k, tCrA_q) LDSM_N 加载 s_acc → B 格式(tCrB) |
| Phase 2 | LDSM_N 加载 v → C 格式(v_bf16) LDSM_N 加载 INV → A 格式(tCrA_k) |
| Phase 4 | LDSM_N 加载 Mqk → A 格式(tCrA_k) |
| Phase 6 | LDSM_T 加载 kr^T → A 格式(ring_A_kr) LDSM_T 加载 s_acc_T → C 格式(ring_S_acc) |
1.4、从寄存器到 smem:STSM 指令
STSM = Store Shared Memory(stmatrix PTX 指令)
和 LDSM 相反,32 线程协作把寄存器中的数据写回 smem。
| Phase | STSM 操作 |
|---|---|
| Phase 5 | STSM_N:out_bf16(C 格式)→ smem out_tile |
| Phase 6 | STSM_T:ring_S_acc(C 格式)→ smem s_acc_T(转置写回) |
1.5、格式转换:为什么需要、怎么做
1.5.1 问题场景
Phase 3 的情况:
u_acc[i] = k_decayed @ s_acc← gemm 输出,C 格式(fp32)u_bf16[i] = bf16(u_acc[i])← 类型转换,还是 C 格式u = (v - u) * beta← 逐元素,还是 C 格式- 现在要算
INV @ u← u 需要作为 B 操作数!!但 u 在 C 格式里!!
C 格式中,lane 5 持有 (row1,col2), (row1,col3), (row9,col2)… B 格式中,lane 5 应该持有同一列的不同行的元素。
完全不同的排列!
1.5.2 传统方案:smem round-trip
1 | C 格式寄存器 ──STSM──> shared memory ──LDSM──> B 格式寄存器 |
- 两次 smem 访问,每次 ~30 cycle 延迟
- 还要
__syncthreads()确保写完再读
1.5.3 MOVM_T:纯寄存器转置
SM75_U32x1_MOVM_T:使用 warp shuffle 在线程间交换寄存器
每个线程把自己的 C 格式寄存器值 shuffle 给需要它的线程,接收方拿到的值自然形成 B 格式。
1 | C 格式的 uint32[0..3] ──4次MOVM_T──> B 格式的 uint32[0..3] |
- 零 smem 访问!
- 延迟只有 shuffle 的 ~5 cycle
代码示例:
1 | uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0)); |
每条 MOVM_T 处理 1 个 uint32 = 2 个 bf16,4 条 = 8 个 bf16 = 一个线程在 16×16 中持有的全部分量。
1.5.4 Phase 中哪里做了格式转换
| Phase | 转换操作 |
|---|---|
| Phase 3(两次,每列块一次) | u_bf16[i] C格式 ──MOVM_T──> tCrB_u_tmp B格式(为了 INV @ u) |
| Phase 4(两次,每列块一次) | u_bf16[i] C格式 ──MOVM_T──> tCrB_u_arr[i] B格式(为了 Mqk @ U 和 Phase6 的 kr^T @ U) |
1.6、load格式 vs 计算格式:tCrAi 和 tCrA 的区别
代码中每个 fragment 都有两份:
tCrAi_k(load 版)和tCrA_k(计算版)tCrAi_q(load 版)和tCrA_q(计算版)tCrBi(load 版)和tCrB(计算版)
原因:
- LDSM 指令要求目标寄存器按特定排列(load layout)
- MMA 指令要求操作数寄存器按另一种排列(compute layout)
这两种 layout 在 SM90 上实际是相同的(底层寄存器完全一致),但 CuTe 的类型系统不知道这一点,它把两种 layout 作为不同类型处理。
所以:
1 | copy(smem → tCrAi_k) // LDSM 写入 load 格式 |
transform + identity 是编译期的类型仪式,不生成实际指令。两个变量名指向的是逻辑上同一组寄存器(或编译器内联后等价)。
为什么不合并成一个?
CuTe 的 copy 和 gemm 要求不同的类型签名。如果直接用 tCrA_k 做 LDSM 目标,类型不匹配会编译报错。tCrAi_k + retile_D 视图让 copy 指令能写入,然后 transform 让 gemm 指令能读取。这是模板元编程的代价。
1.7、累加器(AccFragT)的复用
u_acc[2] 和 out_acc[2] 是 fp32 累加器,每线程持有 2×4 = 8 个 fp32。
为什么 fp32 而不是 bf16?
- MMA 指令的 C 操作数是 fp32
- 多步累加(Phase 1 的 8 步 K 循环)需要高精度防止误差积累
- 只在需要时(Phase 2/3 结尾)才转成 bf16
复用逻辑
u_acc 在 3 个 Phase 中承担不同角色:
| Phase | u_acc 的角色 |
|---|---|
| Phase 1 | u_acc[i] += kd[:, k] @ s_acc[k, warp列i](8步累加,含义:k@s) ↓ Phase 2:值被提取到 u_bf16,u_acc 不再需要 |
| Phase 3 | clear(u_acc[i]) gemm(INV, u_B, u_acc[i])(1步,含义:INV@u = U) ↓ 值被提取到 u_bf16,u_acc 不再需要 |
| Phase 6 | clear(u_acc[bi]) gemm(kr_A, U_B, u_acc[bi])(每步m,含义:kr^T@U) ↓ 值被用于逐元素更新 s_acc |
每次复用前都 clear() 清零。寄存器是宝贵资源,能复用就复用。
out_acc 在 2 个 Phase 中复用:
| Phase | out_acc 的角色 |
|---|---|
| Phase 1 | out_acc[i] += qd[:, k] @ s_acc[k, warp列i](8步累加,含义:q@s) ↓ Phase 2:值被提取到 out_bf16,out_acc 不再需要 |
| Phase 4 | clear(out_acc[i]) gemm(Mqk, U_B, out_acc[i])(1步,含义:Mqk@U) ↓ 值被提取到 gemm_bf16,加到 out_bf16 上 |
1.8、bf16 fragment(SFragT)的角色
SFragT = bf16 类型的 C 格式 fragment。每线程 8 个 bf16。
| 变量 | 用途 |
|---|---|
u_bf16[2] |
承载逐元素运算的中间结果。 Phase 3:先存 bf16(k@s),再变成 (v-k@s)*β,再变成 bf16(INV@u) = U Phase 4:U 被 MOVM_T 转走后就不需要了 |
out_bf16[2] |
最终输出的载体。 Phase 2:存 bf16(q@s) Phase 4:+= bf16(Mqk@U),变成最终 output Phase 5:STSM 写到 smem |
v_bf16[2] |
从 smem 加载的 v 值,Phase 3 做完 v-k@s 后就不需要了 |
为什么用 bf16 而不保持 fp32?
- 逐元素运算 (v-u)*β 对精度要求不高,bf16 够用
- 后面要做 MOVM_T 转成 B 格式,MMA 的 B 操作数只接受 bf16
- bf16 寄存器只占 fp32 的一半空间,省寄存器
1.9、Phase 1-6 寄存器使用的完整故事
一个 warp 的视角,追踪每个寄存器变量:
1 | ┌─ Phase 1 ──────────────────────────────────────────────────┐ |
1.10、逐元素运算为什么必须在 C 格式下做
Phase 3:u = (v - k@s) * beta
Phase 4:out = out_bf16 + gemm_bf16
Phase 6:s_new = s_old * g + gemm_result
这些逐元素运算都在 C 格式下进行。原因:
- gemm 输出本来就是 C 格式,不需要额外转换
- 同一行的 beta/g 值相同,C 格式中同行元素在同一线程 → 只需一个标量
- 逐元素运算每线程独立,无需跨线程通信
如果在 A 或 B 格式下做:
- 同一行的元素可能分散在不同线程 → 需要 shuffle 传 beta 值
- 格式转换本身就有开销
- 完全没必要
所以设计原则是:逐元素运算在 C 格式做,需要作为 operand 时再 MOVM_T 转格式。
1.11、tCrA_k 的一生:理解寄存器名复用
tCrA_k 这个变量在 Phase 1→2→3→4 中被反复覆盖:
| Phase | tCrA_k 的内容 | 操作 |
|---|---|---|
| Phase 1 循环中 | kd[:, k*16:(k+1)*16] |
每步覆盖,存 k_decayed 的列块 |
| Phase 2 | INV[16x16] |
覆盖,存 INV |
| Phase 3 | 被 gemm 读取 | 使用(INV @ u) |
| Phase 4 | Mqk[16x16] |
覆盖,存 Mqk |
| Phase 4 | 被 gemm 读取 | 使用(Mqk @ U) |
它始终是 A 格式,始终通过 LDSM_N 加载(覆盖旧值)。只是内容从 kd 变成 INV 再变成 Mqk。
为什么能这样?
因为 kd、INV、Mqk 的生命期不重叠。Phase 1 用完 kd 后再也不需要它,Phase 2 就可以覆盖成 INV。Phase 3 用完 INV 后,Phase 4 覆盖成 Mqk。
1.11、为什么 Phase 3 的 MOVM_T 和 Phase 4 的 MOVM_T 是两次
Phase 3:
u_bf16(C格式)──MOVM_T──>tCrB_u_tmp(B格式,临时)- 用途:
INV @ u(u 作为 B 操作数) - 结果:
u_acc = INV @ u→u_bf16 = bf16(u_acc)← u_bf16 被新值覆盖了!
Phase 4:
u_bf16(C格式,新值=U)──MOVM_T──>tCrB_u_arr[i](B格式,持久)- 用途:
Mqk @ U(U 作为 B 操作数) - 额外:Phase 6 还要用
tCrB_u_arr做kr^T @ U
为什么不能一次 MOVM_T 搞定?
因为 Phase 3 的 MOVM_T 转的是 u(还没乘 INV 之前),Phase 4 的 MOVM_T 转的是 U = INV @ u(乘完之后)。它们是不同的矩阵!
1 | Phase 3: u → B格式 → INV@u → 得到 U (C格式) |
中间必须经过一次 gemm(INV@u),gemm 的输出必然是 C 格式,所以必须再做一次 MOVM_T 把新的 C 格式转成 B 格式。
1.13、总结:格式转换路径图
1 | smem kd ──LDSM_N──> [A格式] ──┐ |
2 K2 kernel 逐行讲解 Part 1: Layout + Shared Memory + 初始化**
文件: FlashKDA/csrc/smxx/fwd_kernel2.cuh
行 9-70: Layout 定义
1 | D = 128, CHUNK = 16 |
两个关键 Atom
-
Layout_K_INTER_Atom**<bf16>**:
8行 x 16列的 bf16 swizzle layout -
用于 MMA 的 A/B 操作数 (按 K 维度内部交错)
-
避免 shared memory bank conflict
-
Layout_MN_INTER_Atom**<bf16>**: 转置版,
16行 x 8列 -
用于读写转置方向的数据
tile_to_shape
把小 atom 重复铺满整个 shape:
-
例:
tile_to_shape(8x16 atom, (16, 128))= 把 8x16 块铺成 16x128 -
行方向重复
16/8 = 2次 -
列方向重复
128/16 = 8次
各 Layout 含义
| Layout | 尺寸 | 用途 |
|---|---|---|
MMALayout |
[16 x 128] | k_decayed, q_decayed, k_restored, v, out 的 smem layout |
TransposedMMALayout |
[128 x 16] | k_restored^T 的 smem layout (转置视图) |
StateSmemLayout |
[128 x 128] | s_acc 的 smem layout (按列读取友好) |
TransposedStateSmemLayout |
[128 x 128] | s_acc 的转置视图 (按行读取友好) |
注:StateSmemLayout 和 TransposedStateSmemLayout 指向同一片 smem,物理数据相同,只是访问模式不同。
| Layout | 尺寸 | 用途 |
|---|---|---|
BetaSmemLayout |
[32] | beta 的 smem layout (1D,连续) |
GTotalLayout |
[128] | g_total 的 smem layout (1D,连续) |
LMLayout |
[16 x 16] | INV, Mqk 的 smem layout |
TMAVOLayout / TMAStateSmemLayout / TMALMLayout |
[1, …] | TMA descriptor 需要 3D tensor,第一维是 batch/tile 索引 |
FP32StateSmemLayout |
[128 x 128] | fp32 版本的 state layout (fp32 state 输入输出时的临时 buffer) |
为什么需要 Swizzle?
- Shared memory 有 32 个 bank。如果多线程访问同一 bank,会串行化 (bank conflict)
- Swizzle:对地址做 XOR 运算,把原本落在同一 bank 的访问分散到不同 bank
Layout_K_INTER_Atom内部:8行x16列的 bf16 数据,用 swizzle 保证 MMA 的 LDSM 指令 32 线程同时读取时不会有 bank conflict
理解要点:不需要理解 swizzle 的具体 XOR 公式,只需要知道:用了 swizzle layout 的 smem,MMA 读取时不会有 bank conflict。
行 72-112: Shared Memory 结构体
1 | template <class Layouts, int InputStages, int OutputStages> |
InputStages = 3(输入流水线3级)OutputStages = 2(输出流水线2级)
state_acc — 隐状态
1 | alignas(128) cute::ArrayEngine<BF16, cosize_v<StateSmemLayout>> state_acc; |
- 大小 =
128 * 128 * 2 = 32768 字节 = 32KB alignas(128): 对齐到 128 字节 (TMA 要求)- 这是 K2 里最大的一块 smem,常驻不动,chunk 之间保留
InputStorage — K1 的输出
1 | struct InputStorage { |
- 一个
InputStorage约 17.5KB (含 alignas padding) - 这些都是 K1 的输出 + v + beta
OutputStorage
1 | struct OutputStorage { |
Union: 复用内存
1 | union { |
为什么能共用?
- fp32 state 转换只发生在 主循环开始前 (加载 initial state) 和 主循环结束后 (存 final state)
- pipeline buffer 只在 主循环中 使用
- 两者不会同时需要,所以可以 union
Pipeline Barriers
1 | typename cutlass::PipelineTmaAsync<InputStages>::SharedStorage load_pipeline; |
- pipeline 的 barrier 存在 smem 中 (CUTLASS 要求)
state_acc_tma_barrier: 专门给 initial state 的 TMA 用的 barrier
Shared Memory 布局图
1 | 地址: 0 ~82KB |
行 114-151: Kernel 函数签名
1 | __global__ void __launch_bounds__(NumThreads) |
CUTE_GRID_CONSTANT
- 告诉编译器这些 TMA descriptor 存在 constant memory
- GPU 所有线程共享同一份,只读,有缓存
__launch_bounds__(NumThreads)
- 告诉编译器这个 kernel 最多用
NumThreads个线程 - 编译器据此优化寄存器分配
TMA Descriptor 统计
- 8 个 load + 1 个 load state + 1 个 store state + 1 个 store out = 11 个
- 这些都在 host 端 (
fwd_launch.cu) 创建好,通过 kernel 参数传入
行 152-197: 类型定义 + Warp 角色分配
1 | using BF16 = cutlass::bfloat16_t; |
Warp 角色分配表
| threadIdx.x | 0-31 | 32-63 | 64-95 | 96-127 | 128-159 | 160-191 |
|---|---|---|---|---|---|---|
| warp_id | 0 | 1 | 2 | 3 | 4 | 5 |
| role | MMA | MMA | MMA | MMA | LOAD | STORE |
- 4 个 warp (128 线程) 负责 MMA 计算
- 1 个 warp (32 线程) 负责从 workspace 加载数据
- 1 个 warp (32 线程) 负责存储输出
行 172-213: Transaction Bytes + Pipeline 创建
Transaction Bytes 计算
1 | constexpr uint32_t kTmaTransactionBytes = |
- 这个数字告诉 pipeline barrier:一次要搬这么多字节
- 当 TMA 搬完这么多字节后,barrier 会自动 “满足”,通知 consumer
Shared Memory 初始化
1 | extern __shared__ __align__(128) unsigned char shared_mem[]; |
extern __shared__: kernel launch 时由 host 指定大小的动态 shared memoryreinterpret_cast: 把 raw byte 指针当作SharedStorageK2结构体用
创建 Load Pipeline
1 | LoadPipeline load_pipeline = make_load_pipeline<InputStages>( |
make_load_pipeline 做了什么:
- 初始化 3 个 barrier (在 smem 中),每个对应一个 stage
- 设置每个 barrier 的 expected transaction bytes
- 根据
warp_role决定这个线程是 producer 还是 consumer
| 角色 | 可调用的方法 |
|---|---|
| producer (LOAD warp) | acquire / commit / tail |
| consumer (MMA warps) | wait / release |
创建 Store Pipeline
1 | StorePipeline store_pipeline = make_store_pipeline<OutputStages>( |
行 215-237: 序列信息
1 | int seq_idx = blockIdx.x; // grid 的 x 维: 第几个序列 |
变长模式 vs 固定长度模式
变长模式 (IsVarlen = true):
1 | bos = cu_seqlens[seq_idx]; // 序列起始位置 |
固定长度模式 (IsVarlen = false):
1 | int T_seq = T_total / N; // 每个序列长度相同 |
辅助变量
1 | int seq_len = int(eos - bos); |
elect_one_sync(): warp 内的 32 个线程投票,只有一个返回 true- 用于 TMA:只需要一个线程发指令
tile_base 的作用
用于索引 workspace:K1 的输出按 (head, tile) 排列
1 | ws_idx = head_idx * total_tiles + tile_base + t |
示例: 3个序列,长度分别 48, 32, 64,CHUNK=16
| 序列 | bos | eos | tile_base | t_tiles |
|---|---|---|---|---|
| 0 | 0 | 48 | 0 | 3 |
| 1 | 48 | 80 | 3 | 2 |
| 2 | 80 | 144 | 5 | 4 |
行 240-313: 加载 initial_state
情况 1: bf16 state (HasStateIn && !StateFP32)
1 | if constexpr (HasStateIn && !StateFP32) { |
数据流图:
1 | global memory: initial_state[N*H][128][128] bf16 |
情况 2: fp32 state (HasStateIn && StateFP32)
和 bf16 类似,但多了转换步骤:
1 | // 1. TMA 加载到 state_fp32_buf (64KB,union 空间) |
数据流图:
1 | global memory: initial_state[N*H][128][128] fp32 = 64KB |
情况 3: 无 initial state
1 | BF16* buf = shared_storage.state_acc.begin(); |
192 线程并行清零:
| 线程 | 写入位置 |
|---|---|
| 线程 0 | buf[0], buf[192], buf[384], … |
| 线程 1 | buf[1], buf[193], buf[385], … |
- 每线程写
16384/192 ≈ 85个 bf16 零值
3 K2 kernel 逐行讲解 Part 2: LOAD warp + MMA 准备 + Phase 1
行 316-418: LOAD warp 主循环
1 | __syncthreads(); |
获取全局内存视图
1 | Tensor g_v = tma_load_v.get_tma_tensor(make_shape(H, T_total, D)); |
初始化 Pipeline 状态
1 | LoadPipelineState load_write = cutlass::make_producer_start_state<LoadPipeline>(); |
主循环结构
1 | for (int t = 0; t < t_tiles; ++t) { |
加载 v
1 | auto v_off = g_v.layout()(head_idx, int(bos) + t * CHUNK, 0); |
加载 beta
1 | int beta_linear = head_idx * T_total + (int(bos) + t * CHUNK); |
为什么 load 32 个 beta 而不是 16 个?
- TMA 对 1D tensor 有对齐要求,beta 可能不在 8 的倍数地址上
- 向下对齐后多 load 一些,MMA warp 用
beta_smem_offset跳过前面的
示例:
1 | beta_linear = 35, beta_aligned = 32 |
加载 Workspace (以 k_decayed 为例)
1 | { |
注:q_decayed, k_restored, g_total, INV, Mqk 的加载代码结构完全相同,只是换了不同的 TMA descriptor 和 smem 目标。
Stage 推进
1 | ++load_write; |
每次循环加载的数据量
| 数据 | 大小 |
|---|---|
| v | 16×128×2 = 4096 B |
| beta | 32×2 = 64 B |
| k_decayed | 16×128×2 = 4096 B |
| q_decayed | 16×128×2 = 4096 B |
| k_restored | 16×128×2 = 4096 B |
| g_total | 128×4 = 512 B |
| INV | 16×16×2 = 512 B |
| Mqk | 16×16×2 = 512 B |
| 总计 | 17984 字节 ≈ 17.6 KB/chunk |
LOAD warp 到此结束,不参与后续计算。
行 422-428: MMA warp 主循环开始
1 | if (warp_role == WarpRole::MMA) { |
主循环
1 | for (int t = 0; t < t_tiles; ++t) { |
此时:smem 中的 input[load_stage] 已经有了当前 chunk 的所有数据。
时序:
- LOAD warp 已经搬完 chunk t 的数据到
input[load_stage] - MMA warp 开始用
input[load_stage]计算 - 同时 LOAD warp 可能在搬 chunk t+1 到另一个 stage
行 441-527: 取出 smem tensor + 准备 fragment
创建 Smem Tensor 视图
1 | Tensor v_tile = make_tensor( |
Warp/线程信息
1 | const int warp_id = compute_tid / 32; // 0,1,2,3 |
128 个 MMA 线程的职责
| Warp | 线程范围 | 负责列 | 列块数 |
|---|---|---|---|
| warp0 | 0-31 | 0-31 | 2 (0-15, 16-31) |
| warp1 | 32-63 | 32-63 | 2 (32-47, 48-63) |
| warp2 | 64-95 | 64-95 | 2 (64-79, 80-95) |
| warp3 | 96-127 | 96-127 | 2 (96-111, 112-127) |
每个 Warp 内的 group_id 到行的映射 (MMA m16n8k16)
| group_id | lane 范围 | 行 (第一块) | 行 (第二块) |
|---|---|---|---|
| 0 | 0-3 | 0,1 | 8,9 |
| 1 | 4-7 | 2,3 | 10,11 |
| 2 | 8-11 | 4,5 | 12,13 |
| 3 | 12-15 | 6,7 | 14,15 |
| 4 | 16-19 | 8,9 | 0,1 |
| 5 | 20-23 | 10,11 | 2,3 |
| 6 | 24-27 | 12,13 | 4,5 |
| 7 | 28-31 | 14,15 | 6,7 |
Copy 对象
1 | // A 操作数的 copy: smem (K_INTER layout) -> 寄存器 (LDSM_N) |
Copy 对象总结
| 用途 | 对象 | 指令 | 方向 | 使用场景 |
|---|---|---|---|---|
| 加载 A | smem_thr_copy_A |
LDSM_N | smem → reg | Phase 1,2,3,4 (kd, qd, INV, Mqk) |
| 加载 A 转置 | smem_thr_copy_A_T |
LDSM_T | smem → reg | Phase 6 (k_restored^T) |
| 加载 B | smem_thr_copy_B |
LDSM_N | smem → reg | Phase 1 (s_acc) |
| 读 C | smem_thr_load_C |
LDSM_N | smem → reg | Phase 2 (v) |
| 写 C | smem_thr_store_C |
STSM_N | reg → smem | Phase 5 (out) |
| 读 C 转置 | smem_thr_load_C_T |
LDSM_T | smem → reg | Phase 6 (s_acc_T) |
| 写 C 转置 | smem_thr_store_C_T |
STSM_T | reg → smem | Phase 6 (s_acc_T) |
创建参考 Tensor 和 Fragment
1 | Tensor A_ref = local_tile(k_decayed, make_shape(Int<16>{}, Int<16>{}), |
注:这三个 ref 只是用来推导 fragment 的形状,不会真的读这些数据。MMA 一次算 A[16×16] @ B[16×16] → C[16×16],所以参考都是 16×16。
1 | // k_decayed 的 A fragment: 两份 |
thr_mma.partition_fragment_A(A_ref):
- 告诉 CuTe:我这个线程在 MMA 中需要 A 操作数的哪些元素
- 返回一个寄存器 tensor,形状由 MMA atom 决定
- 对于 m16n8k16 + tile 16×16×16:每线程 4 个 uint32 = 8 个 bf16
命名规则:
tCr= thread-level C-format registerAi= A input (load 目标)_k= 给 k_decayed 用_view= retile 后的视图
1 | // q_decayed 的 A fragment (同结构) |
为什么 A 有两份 (k 和 q) 但 B 只有一份?
Phase 1 中 k_decayed 和 q_decayed 用不同的 A,但共享同一个 B (s_acc 的列块)。
Fragment 结构(一个线程):
tCrA_k/tCrA_q:8 个 bf16 (4 个 uint32)tCrB:8 个 bf16 (4 个 uint32)
初始化累加器
1 | AccFragT u_acc[2], out_acc[2]; |
make_fragment_C:创建 MMA 累加器 fragment,fp32- 每线程持有 8 个 fp32 (16×16 矩阵中的 8 个元素)
为什么 [2]?
每个 warp 负责 32 列 = 2 个 16×16 块:
[0]= 前 16 列[1]= 后 16 列
| Warp | 负责列 | [0] 列 |
[1] 列 |
|---|---|---|---|
| warp0 | 0-31 | 0-15 | 16-31 |
| warp1 | 32-63 | 32-47 | 48-63 |
| warp2 | 64-95 | 64-79 | 80-95 |
| warp3 | 96-127 | 96-111 | 112-127 |
行 529-564: Phase 1 — 双 GEMM k@s 和 q@s
目标:
u_acc[0..1]= k_decayed[16×128] @ s_acc[128×(warp的32列)]out_acc[0..1]= q_decayed[16×128] @ s_acc[128×(warp的32列)]
K_BLOCKS = 128 / 16 = 8 步
预取第 0 步
1 | constexpr int K_BLOCKS = decltype(cute::size<1>(k_decayed))::value / 16; |
local_tile(s_acc, (16,16), make_coord(warp_id*2, 0)) 的含义:
s_acc[128×128]按 16×16 分块make_coord(warp_id*2, 0):warp 的第 0 列块,K 维度第 0 步
| Warp | 列块 | s_acc 区域 |
|---|---|---|
| warp0 | 0 | s_acc[0:16, 0:16] |
| warp1 | 2 | s_acc[0:16, 32:48] |
| warp2 | 4 | s_acc[0:16, 64:80] |
| warp3 | 6 | s_acc[0:16, 96:112] |
预取图解
1 | k_decayed [16 x 128] s_acc [128 x 128]: |
主循环实现
1 |
|
Phase 1 指令级并行
每步操作:
transform(零开销)- 预取 B1,k (LDSM)
gemm使用 B0,k (2 条 mma)transformB1,k- 预取 A_{k+1} 和 B0,k+1 (如需)
gemm使用 B1,k (2 条 mma)
每步:4 条 gemm = 8 条 mma 指令 8 步总计:64 条 mma 指令/warp
Phase 1 结束后的寄存器状态
| 累加器 | 内容 |
|---|---|
u_acc[0] |
k_decayed @ s_acc[:, warp第0列块] [16×16] fp32 |
u_acc[1] |
k_decayed @ s_acc[:, warp第1列块] [16×16] fp32 |
out_acc[0] |
q_decayed @ s_acc[:, warp第0列块] [16×16] fp32 |
out_acc[1] |
q_decayed @ s_acc[:, warp第1列块] [16×16] fp32 |
| ``` |
4 K2 kernel 逐行讲解 Part 3: Phase 2-6 + STORE warp
行 566-583: Phase 2 — 类型转换 + 加载 v/INV/beta
1 | SFragT out_bf16[2]; |
out_acc[i]是 fp32 (MMA 累加器)- 转成 bf16 存到
out_bf16[i] - 截断精度: fp32 的 23 位尾数 → bf16 的 7 位尾数
此时 out_bf16 = bf16(q_decayed @ s_acc),后面 Phase 4 会加上 Mqk @ U。
加载 v
1 | SFragT v_bf16[2]; |
从 smem 加载 v 到寄存器,按 C 格式(因为后面要和 u 做逐元素减法)。
v_tile[16×128] 按 16×16 分块:
| Warp | v_block[0] |
v_block[1] |
|---|---|---|
| warp0 | v[:, 0:16] | v[:, 16:32] |
| warp1 | v[:, 32:48] | v[:, 48:64] |
| warp2 | v[:, 64:80] | v[:, 80:96] |
| warp3 | v[:, 96:112] | v[:, 112:128] |
加载 INV
1 | copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(INV), |
- 加载 INV [16×16] 到 A fragment
- 所有 4 个 warp 加载同一份 INV(它只有 16×16,不分列块)
- Phase 3 用它做
U = INV @ u
加载 beta 并计算 sigmoid
1 | BF16 beta0 = BF16(sigmoid_tanh_approx_f32( |
-
group_id = (lane_id / 4) % 8,范围 0-7 -
MMA 的 16 行分成两组:
-
行 0-7:由
group_id0-7 的线程负责,用beta0 = sigmoid(beta[group_id]) -
行 8-15:同样的线程负责,用
beta1 = sigmoid(beta[group_id+8]) -
每个线程持有的 8 个 C 元素中:
-
4 个属于行 0-7(用 beta0)
-
4 个属于行 8-15(用 beta1)
beta_smem_offset 示例:
1 | beta_smem_offset=3, group_id=2 |
行 585-619: Phase 3 — u = (v - k@s) * beta; U = INV @ u
1 | SFragT u_bf16[2]; |
步骤 1: fp32 → bf16 转换
1 | cute::transform(u_acc[i], u_bf16[i], |
此时 u_bf16[i] = bf16(k_decayed @ s_acc),即 k@s 的列块 i。
步骤 2: u = (v - k@s) * beta
1 |
|
MMA C fragment 的坐标系统:
-
make_coord(make_coord(a, row_half), 0, d) -
a = 0,1:两个 “MMA 重复”(16行被分成 2×8) -
row_half = 0:行 0-7(用 beta0) -
row_half = 1:行 8-15(用 beta1) -
d = 0,1:两个列(每线程在每个半行中持有 2 个值)
每线程持有:2(a) × 2(d) × 2(row_half) = 8 个元素
逐元素含义:
1 | u[i][j] = (v[i][j] - k_decayed@s[i][j]) × sigmoid(beta[i]) |
新值 v 减去旧预测 k@s,乘以写入强度 beta。
步骤 3: MOVM_T 转置(C 格式 → B 格式)
1 | uint32_t* u_c = reinterpret_cast<uint32_t*>(&u_bf16[i](0)); |
u_bf16在寄存器中是 C 格式(MMA 输出格式)- 下面要做
INV @ u,u 需要作为 B 操作数(B 格式)
MOVM_T:
- 把 4 个 uint32 从 C 格式重排成 B 格式
- 底层:warp 内 32 线程互相 shuffle 寄存器
- 每条 MOVM_T 处理 1 个 uint32(= 2 个 bf16)
- 4 条 = 8 个 bf16 = 一个线程在 16×16 中持有的全部 B 分量
传统做法 vs MOVM_T:
- 传统:C → smem (STSM) → smem → B (LDSM):两次 smem 访问
- MOVM_T:C → warp shuffle → B:零 smem 访问!
步骤 4: 创建 B fragment
1 | auto tCrB_u_tmp = thr_mma.partition_fragment_B(B_ref); |
创建一个 B fragment,把 MOVM_T 转置后的值填进去。tCrB_u_tmp 现在持有 u 的 B 格式。
步骤 5: U = INV @ u
1 | clear(u_acc[i]); |
tCrA_k= INV [16×16](Phase 2 加载的)tCrB_u_tmp= u [16×16](B 格式)- 只需要 1 次 gemm(K=16,一步就够),内部 2 条 mma 指令
步骤 6: fp32 → bf16 转换
1 | cute::transform(u_acc[i], u_bf16[i], |
Phase 3 结束后:
u_bf16[0]:U 的 warp 第 0 列块 [16×16] bf16u_bf16[1]:U 的 warp 第 1 列块 [16×16] bf16
行 621-646: Phase 4 — out = q@s + Mqk@U
1 | copy(smem_tiled_copy_A, smem_thr_copy_A.partition_S(Mqk), |
加载 Mqk [16×16] 到 A fragment,覆盖之前的 INV。所有 warp 加载同一份 Mqk。
关键变量声明
1 | BFragT_u tCrB_u_arr[2]; // 保留到 Phase 6! |
tCrB_u_arr 会存 U 的 B 格式,Phase 6 更新 state 时还要用。
再次 MOVM_T:U → B 格式
1 |
|
为什么又做一次 MOVM_T?
Phase 3 结束时 u_bf16[i] 被更新成了 INV@u 的 C 格式结果。现在要把这个新的 U(C 格式)再次转成 B 格式。
注意:这次存到 tCrB_u_arr[i] 而不是临时变量,因为 Phase 6 还要用。
out += Mqk @ U
1 | clear(out_acc[i]); |
out_acc[i] = Mqk @ UtCrA_k= Mqk [16×16],tCrB_u_arr[i]= U [16×16] B 格式gemm_bf16 = bf16(Mqk @ U)out_bf16[i] = out_bf16[i] + gemm_bf16 = q_decayed@s + Mqk@U= 最终输出
行 648-653: Phase 5 — 存储 output
1 |
|
STSM:寄存器 out_bf16[i] → smem out_tile 的对应 16×16 块
4 个 warp 并行写 out_tile[16×128]:
| Warp | 列范围 |
|---|---|
| warp0 | 0-31 |
| warp1 | 32-63 |
| warp2 | 64-95 |
| warp3 | 96-127 |
写完后,STORE warp 会通过 TMA 搬到 global memory。
行 655-727: Phase 6 — 隐状态更新
1 | // s_acc[D, D] = s_acc * g_total + k_restored_t[D, 16] @ U[16, D] |
创建 Phase 6 的 fragment
1 | Tensor tCrAi_kr = make_fragment_like<BF16>( |
Ring buffer 结构(PREFETCH=1,就是单个 buffer):
| Buffer | 内容 |
|---|---|
ring_A_kr[0] |
k_restored^T 的当前行块 |
ring_S_acc[bi][0] |
s_acc^T 的当前行块(bi=0:warp第0列块,bi=1:第1列块) |
ring_g0[0] |
g_total 的行 0-7 对应值 |
ring_g1[0] |
g_total 的行 8-15 对应值 |
预取第 0 行块
1 |
|
k_restored_t [128×16] 按 16×16 分块,第 0 行块 = [0:16, 0:16]。LDSM_T:从 smem 转置加载到寄存器。
1 |
|
s_acc_T [128×128] 按 16×16 分块:make_coord(0, warp_id*2+bi) = 第 0 行块,warp 的第 bi 个列块。LDSM_T:转置加载,让数据在寄存器中按 C 格式排列。
1 | ring_g0[0] = g_total(0 * 16 + group_id); |
g_total[128]是 fp32 标量数组ring_g0[0]=g_total[group_id],对应行 0-7(每个线程不同的行)ring_g1[0]=g_total[group_id+8],对应行 8-15
这些 g_total 值已经是 exp2(cumsum) 形式,直接当衰减因子用。
为什么用 s_acc_T(转置视图)?
s_acc的StateSmemLayout是为 Phase 1 的 B 操作数设计的(按列读取)- Phase 6 需要按行更新,用转置视图 + LDSM_T/STSM_T 实现
- 物理内存不变,只是读写模式变了
Phase 6 主循环
1 |
|
步骤 1: k_restored^T @ U
1 |
|
ring_A_kr= k_restored^T 的第 m 行块 [16×16],A 格式tCrB_u_arr[bi]= U 的第 bi 列块 [16×16],B 格式(从 Phase 4 保留!)u_acc[bi] = k_restored^T[m行块] @ U[warp列块bi][16×16] fp32- 2 次 gemm = 4 条 mma 指令
MOVM_T 的价值体现:U 一直在寄存器里,被 3 次复用:
- Phase 3:INV @ u(作为 B)
- Phase 4:Mqk @ U(作为 B)
- Phase 6:k_restored^T @ U(作为 B)
步骤 2: 预取下一行块(与 gemm 重叠)
1 | if (m + PREFETCH < S_M_BLOCKS) { |
gemm 在算当前行块时,LDSM 在加载下一行块(指令级并行)。
步骤 3: 逐元素更新 s_acc
1 |
|
每个元素的更新公式:
1 | s_new = bf16( f32(s_old) × g + gemm_result ) |
s_old:bf16,从ring_S_acc读取bf16_to_f32:转 fp32(无精度损失,bf16 是 fp32 的子集)× g:fp32 乘法,g 是exp2(g_total[对应行])+ u_acc:fp32 加法,gemm 结果bf16():转回 bf16 存储(截断精度)
步骤 4: 写回 s_acc + 预取下一块
1 | Tensor s_block = local_tile(s_acc_T, |
- STSM_T:转置写回 smem,
ring_S_acc(C 格式寄存器)→s_acc_T的[m, warp列块bi]块 - LDSM_T:预取
s_acc_T的下一个行块(与 STSM_T 写回不冲突,写的是 m,读的是 m+1)
Phase 6 结束:s_acc[128×128] 已全部更新。
Phase 6 数据流图
1 | k_restored^T U (寄存器) s_acc (smem) g_total |
4 个 Warp 的分工(Phase 6)
| m | 所有 Warp 更新 | warp0 更新 | warp1 更新 | warp2 更新 | warp3 更新 |
|---|---|---|---|---|---|
| 0 | s_acc[0:16, :] | [0:16, 0:32] | [0:16, 32:64] | [0:16, 64:96] | [0:16, 96:128] |
| 1 | s_acc[16:32, :] | [16:32, 0:32] | [16:32, 32:64] | [16:32, 64:96] | [16:32, 96:128] |
| … | … | … | … | … | … |
| 7 | s_acc[112:128, :] | [112:128, 0:32] | [112:128, 32:64] | [112:128, 64:96] | [112:128, 96:128] |
行 729-737: 同步 + Pipeline 推进
1 | } |
128 个 MMA 线程的 barrier,确保 4 个 warp 都写完了:
out_tile(Phase 5)s_acc(Phase 6)
才能进入下一个 chunk。
1 | cutlass::arch::fence_view_async_shared(); |
fence_view_async_shared:内存栅栏,STSM 是异步写,fence 确保写入对其他 warp 可见producer_commit(out_write):通知 STORE warp “output[out_stage] 写好了,你可以存了”consumer_release(load_read):通知 LOAD warp “input[load_stage] 我用完了,你可以覆盖了”++load_read:下一个 chunk 用下一个 input stage++out_write:下一个 chunk 写下一个 output stage
行 742-798: STORE warp 主循环
1 | if (warp_role == WarpRole::STORE && lane_predicate) { |
尾部 chunk 处理
1 | BF16* out_stage_ptr = shared_storage.output[stage].out.begin(); |
为什么尾部不用 TMA?
- TMA 会写满整个 16×128 = 2048 个 bf16
- 如果序列只剩 5 个 token,TMA 会越界写 11 行,覆盖下一个序列的数据
- 所以用 raw pointer 手动写,只写
actual_len行
完整 chunk 处理
1 | } else { |
TMA store:smem → global memory,一条指令
tma_store_arrive():发起 TMA store 后标记tma_store_wait<0>():等 TMA store 完成(0 表示等所有 pending store)consumer_release:通知 MMA warp “output[stage] 我存完了,你可以覆盖了”
序列长度 50,CHUNK=16 示例
| chunk | actual_len | 存储方式 | 写入行 |
|---|---|---|---|
| 0 | 16 | TMA store | 0-15 |
| 1 | 16 | TMA store | 16-31 |
| 2 | 16 | TMA store | 32-47 |
| 3 | 2 | 手动写 | 48-49 |
行 782-833: 存储 final_state
所有 chunk 处理完后。
bf16 state(HasStateOut && !StateFP32)
1 | if constexpr (HasStateOut && !StateFP32) { |
STORE warp:TMA 直接从 state_acc 存到 global memory。state_acc[128×128] bf16 = 32KB,一次 TMA store。
fp32 state(HasStateOut && StateFP32)
1 | if constexpr (HasStateOut && StateFP32) { |
fp32 state 输出流程:
- Pipeline 结束,此时 union 空间的 input/output buffer 不再使用
- 全线程把
state_acc(bf16,32KB)转成state_fp32_buf(fp32,64KB)
- fp32 buffer 占 union 空间,正好够用
- STORE warp 通过 TMA 存到 global memory
最后 __syncthreads:确保 TMA store 发起后再退出 kernel(TMA 是异步的,但 kernel 退出前必须保证完成)。
附录(一)
FlashKDA K2 流水线机制详解
一、为什么需要流水线
K2 kernel 的工作: 串行遍历一个序列的所有 chunk,做 delta rule 递推。
朴素做法 (无流水线):
1 | for 每个 chunk: |
总延迟 = t_tiles × (500 + 2000 + 500) = t_tiles × 3000 cycle 其中加载和存储时,MMA 完全空闲
流水线做法: 让加载、计算、存储同时进行:
1 | 时间步: 0 1 2 3 4 |
总延迟 ≈ t_tiles × max(L, C, S) + 启动/排空开销 如果 C 是瓶颈: 总延迟 ≈ t_tiles × 2000 (加载/存储被完全隐藏)
二、三种 Warp 角色
K2 用 warp specialization: 不同 warp 永久承担不同角色。
1 | 192 个线程 = 6 个 warp: |
代码 (utils.cuh 行 79-84):
1 | enum class WarpRole { MMA, LOAD_QKG, STORE, NonParticipant }; |
分配逻辑 (fwd_kernel2.cuh 行 189-197):
1 | warp_id < 4 → MMA |
为什么 warp specialization 而不是全部线程一起干?
- TMA 指令只需 1 个线程发起,硬件自动搬运,多线程没意义
- MMA 计算需要 4 个 warp (128 线程) 才能充分利用 Tensor Core
- 分开后三者可以真正并行,不需要 __syncthreads() 全局同步
三、两条 Pipeline 的类型
K2 使用 CUTLASS 提供的两种 pipeline:
3.1 Input Pipeline: PipelineTmaAsync<3>
- 类型:
cutlass::PipelineTmaAsync<InputStages>(InputStages=3) - 本质: 基于 TMA barrier 的异步 pipeline
- 特点: producer 发 TMA 后不用管,硬件搬完自动通知 consumer
- stages: 3 (三缓冲)
角色分配:
- LOAD warp → Producer (发 TMA,绑 barrier)
- MMA warps → Consumer (等 barrier,读数据)
创建 (utils.cuh 行 88-116):
1 | make_load_pipeline<3>( |
关键参数: transaction_bytes = 17984 pipeline 内部的 barrier 会计数: TMA 搬了多少字节? 搬够 17984 字节 → barrier 自动 arrive → consumer 被唤醒
3.2 Output Pipeline: PipelineAsync<2>
- 类型:
cutlass::PipelineAsync<OutputStages>(OutputStages=2) - 本质: 基于 arrive/wait 的软件 pipeline (非 TMA)
- 特点: producer (MMA) 手动 commit,consumer (STORE) 手动 wait
- stages: 2 (双缓冲)
角色分配:
- MMA warps → Producer (STSM 写 output,然后 commit)
- STORE warp → Consumer (wait,然后 TMA store)
创建 (utils.cuh 行 118-143):
1 | make_store_pipeline<2>( |
注意: producer_arv_count=128 128 个 MMA 线程都要 arrive,barrier 才算满足。 这保证 4 个 warp 全部写完 output 后 STORE 才开始搬。
3.3 两种 pipeline 的区别
| PipelineTmaAsync | PipelineAsync | |
|---|---|---|
| 同步机制 | TMA barrier (硬件) | arrive/wait barrier (软件) |
| producer 通知 | TMA 硬件自动 arrive | 手动调 producer_commit |
| consumer 通知 | 自动 (搬完字节数达标) | 手动等 arrive 计数达标 |
| 用途 | global→smem (TMA load) | smem→global (TMA store) |
| 为什么不同 | TMA load 有硬件字节计数 | STSM 是软件写,没有字节计数 |
四、PipelineState: 状态机
每个 pipeline 的 producer 和 consumer 各有一个 PipelineState:
1 | PipelineState<Stages> { |
index()方法: 返回 stage 编号,即 smem buffer 的下标++操作: index 循环递增 (0→1→2→0→1→2→…),phase 在溢出时翻转
Load pipeline (3-stage):
load_write: LOAD warp 的写指针,指向下一个要写入的 stageload_read: MMA warps 的读指针,指向下一个要消费的 stage
Store pipeline (2-stage):
out_write: MMA warps 的写指针out_read: STORE warp 的读指针
初始化:
1 | load_write = make_producer_start_state<LoadPipeline>() |
五、Pipeline API: 每个调用的含义
5.1 LOAD warp 的 API
1 | // ┌────────────────────────────────────────────────────────┐ |
5.2 MMA warps 的 API (consumer of input, producer of output)
1 | // ── 作为 input consumer ── |
5.3 STORE warp 的 API
1 | // ┌────────────────────────────────────────────────────────┐ |
六、三个 warp 角色的完整代码结构
6.1 LOAD warp (fwd_kernel2.cuh 行 319-417)
1 | if (warp_role == LOAD_QKG && lane_predicate) { |
6.2 MMA warps (fwd_kernel2.cuh 行 421-738)
1 | if (warp_role == MMA) { |
6.3 STORE warp (fwd_kernel2.cuh 行 742-780)
1 | if (warp_role == STORE && lane_predicate) { |
七、Pipeline 时序: 详细展开 (t_tiles=5)
假设: LOAD 约 1 单位时间,MMA 约 2 单位,STORE 约 1 单位
1 | 时间: 0 1 2 3 4 5 6 7 8 9 10 11 12 |
stage 使用情况 (input pipeline):
1 | t=0: L→stg0 stg0: 被写入 |
7.1 启动阶段 (pipeline filling)
- t=0: LOAD 写 stg0,MMA 在等 (consumer_wait 阻塞)
- t=1: LOAD 写 stg1,stg0 barrier arrive → MMA 开始 C0
- t=2: LOAD 写 stg2,MMA 还在算 C0
- t=3: 3 个 stage 全满! LOAD 的 acquire 阻塞,等 MMA release
这就是三缓冲的价值: LOAD 连发 3 个 TMA 不停顿。 如果 LOAD 比 MMA 快,这 3 步的加载延迟被完全隐藏。
7.2 稳态阶段 (pipeline steady state)
每当 MMA release 一个 stage:
- LOAD acquire 通过 → 写新数据 → ++stage
- MMA consumer_wait 下一个 stage → 数据已就绪 (提前搬好的)
稳态下: MMA 是瓶颈,LOAD 和 STORE 的延迟被完全隐藏。 总延迟 ≈ t_tiles × C (MMA时间) + 启动排空开销
7.3 排空阶段 (pipeline draining)
- LOAD 搬完最后一个 chunk → producer_tail()
- MMA 算完最后一个 chunk → 不再 consumer_wait
- STORE 存完最后一个 chunk → 循环结束
八、两条 Pipeline 的交互
MMA warps 同时是 input pipeline 的 consumer 和 output pipeline 的 producer。
每个 chunk 的 MMA 循环体开头要 两个等待:
1 | store_pipeline.producer_acquire(out_write); ← 等 output stage 可写 |
为什么 store acquire 在 load wait 前面? 先确认 output buffer 可写,再等输入数据。 如果反过来: 数据到了但 output buffer 没空,还是得等。 实际上两个 wait 可以任意顺序,但先检查更快满足的那个可以微微减少等待。
循环体结尾:
1 | compute_barrier.arrive_and_wait(); ← 4 个 warp 同步 |
两个通知同时发出,STORE 和 LOAD 同时被唤醒。
九、Barrier 的底层机制
9.1 TMA Barrier (Input Pipeline)
Hopper (SM90) 新增的硬件 barrier,存在 shared memory 中。
核心能力: 字节计数。
arrive_and_expect_tx(N): 告诉 barrier “我期望 N 字节”- TMA 硬件每搬完一块: barrier.bytes_arrived += 搬运字节数
- 当 bytes_arrived >= expected → barrier 自动满足
1 | // LOAD warp: |
这就是为什么 LOAD warp 不需要显式 commit: TMA 硬件自己在计数,搬完就通知。
9.2 Software Barrier (Output Pipeline)
Output pipeline 用的是软件 arrive barrier:
1 | producer_commit: 128 个 MMA 线程各调一次 arrive() |
为什么不用 TMA barrier? 因为 output 是 MMA warp 用 STSM 写入 smem 的 (不是 TMA)。 STSM 是普通的 smem 写指令,没有硬件字节计数能力。 所以用软件 arrive 计数代替。
十、compute_barrier vs pipeline barrier
两种不同的 barrier,容易混淆:
1 | // ┌─────────────────────────────────────────────────────────────────┐ |
时间线:
1 | consumer_wait(input) ← 等 LOAD (跨角色) |
十一、fence_view_async_shared 的作用
代码 (行 732):
1 | cutlass::arch::fence_view_async_shared(); |
放在 compute_barrier 之后,producer_commit 之前。
为什么需要: Phase 5 用 STSM 写 out_tile (寄存器→smem) Phase 6 用 STSM_T 写 s_acc (寄存器→smem)
STSM 是异步指令: 发出后不保证立刻写到 smem 的全局可见视图。 compute_barrier 只保证本 warp 的 STSM 已发出,不保证其他 warp 可见。
fence_view_async_shared:
确保当前线程之前的所有 STSM 对其他 warp/线程可见。
STORE warp 读 smem 时一定能看到最新的数据。
顺序:
1 | STSM (Phase 5, 6) |
十二、Shared Memory 的生命周期管理
1 | ┌─────────────────────────────────────────────────────────────┐ |
十三、完整 API 调用时序图
以一个 chunk (t=2) 为例,展示所有 pipeline API 调用:
1 | LOAD warp (warp4) MMA warps (warp0-3) STORE warp (warp5) |
附录(二)
一、为什么需要 Prefetch
GPU 上两类操作的延迟差异:
- LDSM (smem → 寄存器): ~20 cycle
- LDSM_T (smem → 寄存器): ~20 cycle
- MMA (矩阵乘): ~16 cycle
- ALU (逐元素运算): ~5-10 cycle
如果串行执行: LDSM 等 20 cycle → MMA 等 16 cycle → LDSM 等 20 cycle → … 大量时间浪费在等待上
Prefetch 的做法: 先发 LDSM 加载"下一步"的数据,然后立刻执行"当前步"的 MMA/ALU。 LDSM 和 MMA/ALU 在硬件上使用不同的功能单元,可以并行执行。 等 MMA 做完,LDSM 的结果也差不多到了。
关键前提: LDSM 用的是 Load/Store 单元,MMA 用的是 Tensor Core, ALU 用的是 CUDA Core,三者互不干扰。
二、K2 中的四层 Prefetch / 缓冲
从离计算最远 (global memory) 到最近 (寄存器),有四层缓冲:
| 级别 | 缓冲深度 | producer | consumer | 每份大小 | 隐藏的延迟 |
|---|---|---|---|---|---|
| TMA input | 3-stage | LOAD warp | MMA warps | ~14 KB(smem) | global→smem (~500c) |
| TMA output | 2-stage | MMA warps | STORE warp | 4 KB(smem) | smem→global (~500c) |
| Phase1 B0/B1 | 2 份 | LDSM | gemm | 8 bf16(reg) | smem→reg (~20c) |
| Phase6 ring | 1 份 | LDSM_T | gemm+ALU | 8-16 bf16 | smem→reg (~20c) |
越靠近 global memory 延迟越大,需要的缓冲深度越深。 越靠近寄存器延迟越小,单缓冲 + 代码排序就够了。
三、TMA 3-stage Input Pipeline —— 三缓冲
3.1 结构
1 | smem input[0]: buffer A ─┐ |
3.2 为什么三缓冲而不是双缓冲
双缓冲只能领先 1 步:
1 | 双缓冲: |
三缓冲多一步余量:
1 | 三缓冲: |
核心收益: LOAD 能连续发 3 个 TMA 不停顿,前几个 chunk 的加载延迟被完全隐藏。
3.3 同步机制
1 | LOAD warp: |
stage 0 的生命周期: L0 写入 → C0 消费 → release → L3 写入 → C3 消费 → release → L6 写入 → …
3.4 Output Pipeline (2-stage) 也是多缓冲
1 | smem output[0], output[1]: 双缓冲 |
只需双缓冲: MMA 写一个 stage 时 STORE 读另一个,足够隐藏 TMA store 延迟。
四、Phase 1 的 Prefetch —— 双列块交替 (寄存器级双缓冲)
4.1 背景
Phase 1: u_acc = kd[16×128] @ s_acc[128×128] (每 warp 取 32 列)
沿 K 维度分 8 步 (K_BLOCKS=8),每步:
- A = kd[:, k*16:(k+1)*16] 16×16
- B = s_acc 的两个列块 (B0, B1) 各 16×16
每步 4 条 gemm (kd@B0, qd@B0, kd@B1, qd@B1)
问题: B fragment 只有一份寄存器 (tCrBi/tCrB),怎么处理 B0 和 B1?
4.2 双缓冲设计
1 | tCrBi = load buffer (LDSM 写入目标) |
两个变量交替: 一个在被 LDSM 写入,一个在被 gemm 读取。 transform(tCrBi → tCrB) 做"交接"。
4.3 完整展开 (以 k=0 为例)
1 | ═══ 循环前: 预取第 0 步的 A 和 B0 ═══ |
4.4 寄存器状态变化表
| tCrAi_k (load) | tCrA_k (计算) | tCrAi_q (load) | tCrA_q (计算) | tCrBi (load) | tCrB (计算) | |
|---|---|---|---|---|---|---|
| 循环前: | kd[K=0] | ─ | qd[K=0] | ─ | B0[K=0] | ─ |
| ↑ 预取 | ↑ 预取 | ↑ 预取 | ||||
| ① trans: | ────→ | kd[K=0] | ────→ | qd[K=0] | ────→ | B0[K=0] |
| 可用gemm | 可用gemm | 可用gemm | ||||
| ② LDSM: | B1[K=0] | |||||
| ↑加载中 | B0仍可用 | |||||
| ③ gemm: | kd @ B0 | qd @ B0 | B0 被读取 | |||
| ④ trans: | ────→ | B1[K=0] | ||||
| 可用gemm | ||||||
| ⑤ LDSM: | kd[K=1] | qd[K=1] | B0[K=1] | |||
| ↑加载中 | ↑加载中 | ↑加载中 | B1仍可用 | |||
| ⑥ gemm: | kd @ B1 | qd @ B1 | B1 被读取 | |||
| ← ⑤加载完成 | 下步①可用 |
4.5 为什么 B 有 B0/B1 交替而 A 没有
每步 k:
- A (kd, qd): 同一个 16×16 块,被 B0 和 B1 两次 gemm 共享
- B (s_acc): 两个不同的 16×16 块 (列块0 和 列块1)
A 在一步内不变,只在 k→k+1 时更新 B 在一步内要换两次 (B0→B1)
所以 B 需要 “用 B0 时预取 B1,用 B1 时预取下步 B0” 而 A 只需要 “用完后预取下步”
4.6 指令级并行的时间线
1 | 时间 → |
五、Phase 6 的 Prefetch —— 单缓冲 + 精确排序
5.1 背景
Phase 6: s_acc = s_acc * g + kr^T @ U
8 步循环 (m=0…7),每步更新 s_acc 的 16 行。 三种数据需要 prefetch:
- ring_A_kr: kr^T 的行块 [16×16] LDSM_T 从 smem
- ring_S_acc: s_acc 的行块 [16×16] LDSM_T 从 smem
- ring_g0/g1: g_total 标量 从 smem 读
5.2 PREFETCH=1 意味着什么
PREFETCH=1 → ring buffer 深度为 1 ring_A_kr[0], ring_S_acc[bi][0], ring_g0[0], ring_g1[0]
slot = m % 1 = 0 (永远是 0)
没有 “环形” 效果,就是单个 buffer,靠代码顺序保证 “先读后写”。 如果 PREFETCH=2,就会有 ring[0] 和 ring[1] 交替,变成真正的双缓冲。
5.3 完整展开 (m=0, m=1 为例)
1 | ═══ 循环前: 预取 m=0 的数据 ═══ |
5.4 为什么 ③ 和 ⑥ 不能合并到一起
如果把 ⑥ 移到 ③ 旁边:
1 | ② gemm 完成 |
ring_S_acc 还没被 ④ 消费,就被 ⑥ 覆盖了。
每个变量的 “可覆盖” 时机:
| 变量 | 消费完毕的时刻 | 最早可覆盖 | 实际预取位置 |
|---|---|---|---|
| ring_A_kr | ② gemm 读完后 | ③ (紧接②后) | ③ ✓ |
| ring_g0/g1 | ① 已拷贝到局部变量 | ③ (随时) | ③ ✓ |
| ring_S_acc | ⑤ STSM_T 写回后 | ⑥ (紧接⑤后) | ⑥ ✓ |
5.5 ③ 也不能移到 ⑥ 旁边
正确性上可以 (ring_A_kr 在 ② 后就不需要了),但丢失重叠机会:
当前 (③在④前面):
1 | 时间 → |
如果 ③ 挪到 ⑥ 旁边:
1 | 时间 → |
量化: LDSM_T ~20 cycle,④逐元素 ~15-20 cycle
- 当前: ③ 的 20c 被 ④ 完全隐藏,免费
- 挪后: 多出 ~20c 裸等,8 步共多 ~160 cycle / chunk
5.6 Phase 6 一步内的完整时间线
1 | 时间 → |
六、Prefetch 与双缓冲/三缓冲的关系
核心思想完全一样: 加载和计算重叠,隐藏延迟。 只是实现方式因场景而异。
6.1 经典双缓冲
1 | buffer A, buffer B |
关键特征: 2 份存储空间,角色轮换,不同 buffer 永远不冲突。
6.2 K2 中各处的对应
1 | ┌─────────────────────────────────────────────────────────────────┐ |
6.3 对比总结
| 经典双缓冲 | Phase1 tCrBi/tCrB | Phase6 PREFETCH=1 | |
|---|---|---|---|
| 缓冲区数量 | 2 | 2 | 1 |
| 交替方式 | A↔B 轮换 | load↔compute 轮换 | 先读后写,原地覆盖 |
| 正确性保证 | 不同buffer不冲突 | transform做交接 | 代码顺序保证 |
| 能隐藏的延迟 | 完整一步加载延迟 | 完整LDSM延迟 | 部分 (靠与ALU重叠) |
| 额外存储开销 | 2× | 2×(两份fragment) | 1× (无额外开销) |
| ``` |
- 标题: KDA源码剖析之FlashKDA(下)
- 作者: 鱿鱼圈
- 创建于 : 2026-06-05 23:59:00
- 更新于 : 2026-06-14 16:01:41
- 链接: https://yuyanqi.com/2026/06/05/KDA源码剖析之FlashKDA(下)/
- 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。