KDA源码剖析之FlashKDA(上)
前置知识:KDA算法公式、triton、cuda、c++
仓库链接:MoonshotAI/FlashKDA: FlashKDA: high-performance Kimi Delta Attention kernels
公式回顾
kernel 1 代码
代码链接:FlashKDA/csrc/smxx/fwd_kernel1.cuh at master · MoonshotAI/FlashKDA
1 | // ===== Launch Kernel 1 (prepare) ===== |
计算流程图




整体剖析
grid: (total_tiles, H)
threads: 256
第一部分: 数据布局定义 (第5-41行)
1 | template <int D, int CHUNK = 16> |
这是一个模板结构体,定义K1 kernel里所有shared memory的排列方式。
D=128 (head维度), CHUNK=16 (每个chunk的token数)。
QKLayout (第7行)
1 | using QKLayout = make_layout(make_shape(16, 128), LayoutRight{}); |
形状: [16, 128] LayoutRight: 行主序(row-major)
- 地址 =
i * 128 + j - 总大小 = 16 × 128 = 2048 个元素
- 用途: 存放 TMA 加载进来的原始 q 和 k
MMALayout (第9-13行)
1 | using MMALayout = tile_to_shape( |
什么是 “swizzle 布局”?
- GPU shared memory 分成 32 个 bank,每 bank 4字节宽
- 如果一个 warp 的 32 个线程同时访问同一 bank → bank conflict → 串行化
- Swizzle 布局通过"打乱"地址映射,让线程尽量访问不同 bank
GMMA::Layout_K_INTER_Atom**<bf16>**: CUTLASS 预定义的 swizzle 模式,专门为 bf16 MMA 指令优化,确保加载操作数时无 bank conflict。
tile_to_shape: 把这个原子 swizzle 模式"铺开"到 [16, 128] 的大小 LayoutLeft: 列优先的铺开方式
简单理解: MMALayout 和 QKLayout 存的是同样 16×128 的数据,但 MMALayout 的地址映射经过特殊打乱,让 MMA 指令读取时不冲突。
其他布局 (第14-20行)
| 布局 | 形状 | 用途 |
|---|---|---|
BetaSmemLayout |
[32], stride=1 | beta 加载 (实际只用16个,TMA要求对齐) |
GTotalLayout |
[128], stride=1 | g_total 存128个float |
LMLayout |
[16,16] swizzle | L 和 Mqk 矩阵 |
TransposedLMLayout |
[16,16] 转置方向 | 转置版本的 LMLayout |
TMA 布局 (第27-40行)
TMA (Tensor Memory Accelerator) 是 Hopper GPU 的硬件单元:
| 传统做法 | TMA做法 |
|---|---|
| 线程发load指令 | 一个线程下一条命令 |
| 通过L1/L2 cache | 硬件DMA直接搬数据 |
| 数据到寄存器 → 写到smem | 直接搬整块到smem |
好处:
- 不占计算线程的时间,异步执行
- 自动处理2D/3D地址计算(给行列坐标,自动算出全局内存地址)
- 自动处理 swizzle (搬运时直接按 MMALayout 排列)
TMA 布局需要比对应的 smem 布局多一个维度(prepend加了一个size=1的维度),这是 CUTLASS TMA 接口的要求。
第二部分: Shared Memory 结构 (第43-83行)
union 核心 (第56-70行)
1 | union { |
C++ 的 union 意味着两个 struct 共享同一块内存!
Phase A: 存原始 q, k, g
| 变量 | 大小 | 类型 |
|---|---|---|
| q | 16×128×2 = 4096 字节 | bf16 |
| k | 16×128×2 = 4096 字节 | bf16 |
| g | 16×128×4 = 8192 字节 | float (cumsum需要fp32) |
| 合计 | 16384 字节 |
Phase B: 存计算结果 (共用同一块内存)
| 变量 | 大小 | 类型 |
|---|---|---|
| k_decayed | 16×128×2 = 4096 字节 | bf16 (MMALayout) |
| q_decayed | 16×128×2 = 4096 字节 | bf16 (MMALayout) |
| k_inv | 16×128×2 = 4096 字节 | bf16 (MMALayout) |
| L | 16×16×2 = 512 字节 | fp16 |
| INV | 16×16×2 = 512 字节 | bf16 |
| Mqk | 16×16×2 = 512 字节 | bf16 |
| 合计 | 13824 字节 |
两个阶段不重叠(Phase A用完才开始Phase B),所以可以共用。省了约 14KB shared memory,让 SM 能放更多 CTA (occupancy 更高)。
其他内存区域
1 | alignas(128) bf16 beta[32]; // 不在union里,两个阶段都要用 |
ClusterTransactionBarrier:
-
TMA 是异步的: thread 0 发命令后立即返回,数据还没到
-
barrier 用来等 TMA 完成:
-
发 TMA 前:
barrier.arrive_and_expect_tx(字节数)— 告诉barrier"将有这么多数据到达" -
等待:
barrier.wait(0)— 阻塞直到所有预期的数据到达
第三部分: kernel 函数签名 (第85-119行)
1 | __global__ void __launch_bounds__(NumThreads, 8) _flash_kda_fwd_prepare(...) |
| 关键字 | 含义 |
|---|---|
__global__ |
CUDA kernel 函数 |
__launch_bounds__(256, 8) |
每个CTA最多256线程,目标每SM同时运行8个CTA → 编译器限制寄存器使用,让SM能塞下更多CTA |
参数列表 (第100-118行)
1 | CUTE_GRID_CONSTANT TmaLoadQ const tma_load_q |
CUTE_GRID_CONSTANT 告诉编译器这些参数放在 GPU 的 constant memory 里。
TMA descriptor 包含: 全局内存地址、形状、stride、swizzle模式。thread 用 descriptor 来发起 TMA 操作,硬件自动完成数据搬运。
| 类型 | 数量 | 用途 |
|---|---|---|
| Load descriptor | 6 | q, k, beta, g, dt_bias |
| Store descriptor | 6 | kd, qd, kr, gt, inv, mqk 到workspace |
标量参数:
scale: 注意力缩放T_total: 总token数H: head数N: 序列数total_tiles: 总chunk数gate_scale: gate下界×log2e
第四部分: 每个CTA确定自己处理哪个chunk (第147-181行)
1 | int global_tile_idx = blockIdx.x; |
blockIdx.x: 全局chunk编号 (0, 1, …, total_tiles-1)blockIdx.y: head编号 (0, 1, …, H-1)
找到当前chunk属于哪个序列
Varlen模式 (多个变长序列拼在一起):
1 | 线性扫描 cu_seqlens 数组: |
非Varlen模式 (所有序列等长):
1 | T_seq = T_total / N |
边界检查
1 | if (local_t >= t_tiles_this_seq) return; |
因为 total_tiles 是上界(varlen时可能多分配了一些),多余的CTA直接退出。
第五部分: TMA加载输入数据 (第182-244行)
thread 0 发起所有TMA
1 | if (threadIdx.x == 0) { |
只有 thread 0 发 TMA 命令。TMA 是硬件操作,一个线程发命令就够了,其他 255 个线程不需要参与。
初始化 barrier
1 | shared_storage.tma_load_barrier.init(1); |
告诉 barrier “接下来会有 kTmaTransactionBytes 字节的数据到来”。
1 | kTmaTransactionBytes = 3×16×128×2 + 32×2 + 128×4 |
这是 q + k + g_bf16 + beta + dt_bias 的总大小。
准备TMA地址
1 | Tensor g_q = tma_load_q.get_tma_tensor(make_shape(H, T_total, D)); |
发起TMA
1 | s_q_tile = make_tensor(smem_ptr(shared_storage.q), TMAQKLayout{}) |
这一行就是 TMA 操作:
- 源(S): 全局内存的 g_q_tile
- 目的(D): shared memory 的 s_q_tile
- with(barrier): 绑定到 barrier,完成时自动通知
TMA 硬件会:
- 根据 descriptor 算出全局内存地址
- 发起 DMA 请求
- 数据到达 shared memory 后,自动给 barrier 的计数器减(已传输字节数)
- 当 barrier 的计数器减到 0,
barrier.wait()解除阻塞
q, k, g_bf16, beta, dt_bias 共 5 次 TMA,都绑定同一个 barrier。
CPU计算与TMA重叠
1 | float a_log_exp = expf(A_log_ptr[head_idx]); |
同步模式 (Hopper 标准)
1 | __syncthreads(); // 确保barrier初始化可见 |
第六部分: L2归一化 (第246-285行)
线程分工
1 | constexpr int ELEMS_PER_THREAD = 8; |
每行 128 维,16 个线程分工,每线程处理 8 个元素。 256 个线程 / 16 = 16 行,正好覆盖 CHUNK=16 行。
1 | int my_row = threadIdx.x / 16; // 我负责第几行 |
示例: threadIdx.x = 35
my_row = 35 / 16 = 2(第2行)my_col = (35 % 16) × 8 = 3 × 8 = 24(第24-31列)
读数据 + 计算平方和
1 | for i in 0..7: |
每个线程累加自己 8 个元素的平方和。
warp shuffle 规约
1 | for delta = 8, 4, 2, 1: |
__shfl_xor_sync: warp 内的寄存器数据交换。
lane i 和 lane i^delta 互换 q_sq 值,然后各自相加。
| delta | 效果 |
|---|---|
| 8 | lane 0-7 和 lane 8-15 交换 → 每个lane有16个元素的平方和 |
| 4 | 相邻4个交换 |
| 2 | 相邻2个交换 |
| 1 | 最终同一行的16个线程都有完整的行平方和 |
为什么从 delta=8 开始而不是 16?
因为 THREADS_PER_ROW=16,同一行占了半个 warp (lane 0-15)。
delta=8 让 lane 0 和 lane 8 交换(同一行内),够了。
归一化
1 | q_inv = rsqrtf(q_sq + 1e-6f) // 1/sqrt(||q||^2 + eps) |
结果直接写回 shared memory (覆盖原来的 q/k)。
这就是 L2 normalization: x = x / ||x||
第七部分: gate激活 + cumsum (第287-319行)
获取实际长度
1 | int actual_len = min(CHUNK, seq_len - local_t * CHUNK); |
当前chunk的实际token数(最后一个chunk可能不满16)。
threads 0-127 做 gate + cumsum (每线程一列)
1 | int col = compute_tid; // 我负责第几列 (0-127) |
gate激活公式:
1 | g_val = gate_scale × sigmoid(A_log_exp × (g_bf16 + dt_bias)) |
其中:
gate_scale = lower_bound × log2(e)≈ -5 × 1.4427 = -7.2135sigmoid用tanh.approx.f32实现:sigmoid(x) = 0.5 × tanh(0.5×x) + 0.5- 最终
g_val的范围是[gate_scale, 0] = [-7.2135, 0] - cumsum 16步后最大绝对值 ≈ 16 × 7.2135 ≈ 115.4
对比fla: fla用单独的 chunk_local_cumsum_vector_kernel 来做 cumsum,并且需要先在外部做 gate 激活。FlashKDA 全部融合在一起。
threads 128-255 清零 k 的 padding 行
1 | for row = actual_len .. 15: |
如果 chunk 不满 16 个 token,把 k 的多余行清零。这样后面矩阵乘时自然不会把 padding 算进去。
第八部分: g_total 转为 exp2 形式 (第334-339行)
1 | if (compute_tid < 128): |
把 g_total 从 cumsum 值变成 exp2(cumsum) 值。后面 decay_apply 需要用 exp2(g_total) 来计算 k_restored。
1 | ex2_approx_ftz_f32`: PTX 指令 `ex2.approx.ftz.f32 |
ex2 = 2^x(base-2 指数)approx= 近似(快速, 约23bit精度)ftz= flush to zero (极小值直接归零)- 比
expf(x)快,因为expf需要先乘log2(e)再调ex2
第九部分: decay_apply (第341-457行)
最复杂的一段。256个线程把 q, k, g 转换成 q_decayed, k_decayed, k_inv, k_restored。
warp/lane 信息
1 | int lane = compute_tid % 32; // warp内的lane编号 (0-31) |
命名 g 和 t 是因为数据访问模式:
g决定读哪一行:row = m_blk + (warp_id + g) % 8t决定读行内的哪 2 个元素: 第t×2和t×2+1个
循环参数
1 | N_M = CHUNK / 8 = 2 // 16行分成2个8行块 |
寄存器数组
1 | float reg_g[4][2] -- 4个tile,每个tile存2个gate值 |
第一轮循环: 读数据到寄存器
1 | for m_blk in {0, 8}: // 两个行块 |
关键: (warp_id + g) % 8 让不同 warp 的不同 lane-group 读不同的行。这样 256 个线程读完了全部 16×128 = 2048 个 (q,k,g) 值。
每个线程在4个tile中各持有2个元素,共 4×2 = 8 个元素。
同步
1 | __syncthreads(); |
同步! 接下来要写到 union 的另一半 (k_decayed/q_decayed/k_inv/k_restored)。必须确保所有线程都读完了 q/k/g,因为写入会覆盖它们。
第二轮循环: 计算并写结果
1 | for m_blk, n_blk (同上): |
注意这里写的是 MMALayout (swizzle布局),不是 QKLayout (行主序)。因为后面 MMA 操作需要 swizzle 布局。
对应fla的关系
| FlashKDA | fla对应 | 公式 |
|---|---|---|
| k_decayed | k × beta × exp2(gk) | = k * exp2(g) |
| q_decayed | qg × scale | = q * exp2(g) * scale |
| k_inv | 用于构造L矩阵 | = k * exp2(-g) |
| k_restored | kg | = k * exp2(g_total - gk) |
第十部分: 构造 L 和 Mqk 矩阵 (第459-468行)
1 | Tensor L = make_tensor(smem_ptr(shared_storage.L), LMLayout{}); |
L 用 fp16 存(和bf16共用同样的 LMLayout 因为都是2字节)。reinterpret_cast 只是改变类型解释,不改变内存。
MMA计算
1 | if (compute_tid < 32) { |
| Warp | 线程 | 计算 | 结果类型 |
|---|---|---|---|
| warp 0 | 0-31 | L = k_decayed @ k_inv^T |
fp16 |
| warp 1 | 32-63 | Mqk = q_decayed @ k_inv^T |
bf16 |
其他 warp 闲着。
mma_m16n16_xxx 内部 (utils.cuh 第146-165行)
- 用
SM80_16x8x16_F32BF16BF16F32_TN的 MMA atom - 输入 bf16,累加 fp32,然后转目标精度存
- 一个 16×16 矩阵乘 = 两个 16×8×16 atom (沿N维拼)
cooperative_gemm让一个warp协作完成
数学: L_ij = sum_d k_decayed[i,d] × k_inv[j,d]
对比fla: fla在token_parallel kernel里逐token循环计算 Akk 的元素;FlashKDA一次MMA直接算出完整的16×16矩阵。
为什么L用fp16而Mqk用bf16? L后面要送进 Neumann 求逆,fp16有10bit尾数(vs bf16的7bit),求逆过程中误差累积更小。Mqk直接用于最终计算,bf16够用。
第十一部分: 下三角化 + INV = I - L (第474-492行)
1 | if (compute_tid < 256): # 所有线程 |
256个线程处理 16×16 = 256 个元素,正好一对一。
对比fla:
- fla在token_parallel kernel里,Akk对角线上是1,
Akk(i,j) = sum_d k[i]×k[j]×exp(g_i-g_j) × beta[j] (j<i) - 然后在inter_solve_fused里做 I - Akk
- FlashKDA这里
L = k_decayed @ k_inv^T已经包含了exp(g_i-g_j)的效果 - 再乘
sigmoid(beta)就完成了
注意: beta的sigmoid是在这里才做的(K1开头TMA加载的是原始logits)。fla里beta是在外部提前sigmoid过的。
第十二部分: Neumann 级数求逆 (第494-498行)
1 | neumann_inv_fused_1warp(L_fp16, INV_fp16, INV, compute_tid); |
只有线程 0-31 (一个warp) 实际计算,其他线程跳过。
算法 (utils.cuh 第189-312行)
1 | L^2 = L x L # MMA |
输入: L (严格下三角, fp16), INV = I - L (fp16)
输出: INV_bf16 = (I-L)^{-1} (bf16)
数学原理
要求 (I + L)^{-1}
Neumann级数: (I + L)^{-1} = I - L + L^2 - L^3 + ... - L^{15}
(L是16×16严格下三角,L^{16} = 0,级数有限项,精确!)
倍增法加速:
1 | INV_0 = I - L |
展开: (I-L)(I+L^2)(I+L^4)(I+L^8) = I - L + L^2 - L^3 + ... - L^{15}
→ 这是 (I+L)^{-1} 的精确展开 (因为 L^{16}=0)
关键优化 — 全在寄存器里算
每个 16×16 MMA 的操作数用 uint32_t[4] 表示 (每个 lane 持有 4 个 fp16 值)
transpose_u32x4: 用 SM75_U32x1_MOVM_T 在寄存器内转置
- A操作数格式 → B操作数格式
- 不需要写回 shared memory!
mma_16x16: 两个 SM80_16x8x16_F16F16F16F16_TN 拼成 16×16
add_fp16x2_u32x4: fp16x2 向量加法 (__hadd2)
最后: 把 fp16 结果转为 bf16 写回 INV smem。
第十三部分: TMA存到workspace (第499-569行)
1 | if (threadIdx.x == 0): |
workspace 按 [H × total_tiles, ...] 排列,ws_idx 是当前tile的全局索引。
6次TMA store
每次的模式相同:
- 构造全局 tensor 描述
- 算出目标偏移
- 构造 smem 侧的 tile
cute::copy(tma_store_xxx, S, D)— 发起TMA storetma_store_arrive()— 通知硬件"有一个store要做"
存储的6个中间结果:
| 输出 | 形状 | 大小 | 含义 |
|---|---|---|---|
| k_decayed | [16,128] bf16 | 4KB | k × exp2(cumsum_gate) |
| q_decayed | [16,128] bf16 | 4KB | q × exp2(cumsum_gate) × scale |
| k_restored | [16,128] bf16 | 4KB | k × exp2(g_total - cumsum_gate) |
| g_total | [128] fp32 | 512B | exp2(chunk总gate) |
| INV | [16,16] bf16 | 512B | (I - L)^{-1} |
| Mqk | [16,16] bf16 | 512B | q_decayed @ k_inv^T |
每个tile总workspace = 3×4096 + 512 + 512 + 512 = 13824 字节 ≈ 13.5KB
等待完成
1 | tma_store_wait<0>(); // 等待所有TMA store完成 (<0>表示等到剩余0个未完成) |
第十四部分: K1 总结
输入
| 输入 | 来源 |
|---|---|
| q, k, g | 全局内存 (TMA加载) |
| beta, dt_bias, A_log | 全局内存 (TMA加载) |
计算流程
1 | 1. L2归一化 q, k (256线程并行) |
输出 (TMA存到workspace)
1 | k_decayed, q_decayed, k_restored, g_total, INV, Mqk |
这些workspace数据会被 K2 kernel 读取,用于chunk间的隐状态递推。
对应fla中的步骤
| FlashKDA K1 | fla对应 |
|---|---|
| 全部融合 | l2norm + gate激活 + cumsum + Step1(token_parallel) + Step2(inter_solve, 矩阵求逆) + Step3(recompute_w_u的部分) |
附录:decay_apply 线程映射详解
文件位置: fwd_kernel1.cuh 的 decay_apply 部分
256个线程从 shared memory 读取 q/k/g,计算 q_decayed/k_decayed/k_inv/k_restored,写回到 shared memory 的 union 复用空间。
一、256个线程的编号规则
1 | int lane = compute_tid % 32; // warp内编号 0-31 |
每个warp内部:
| lane | 0 | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 | 9 | 10 | 11 | … | 28 | 29 | 30 | 31 |
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| g | 0 | 0 | 0 | 0 | 1 | 1 | 1 | 1 | 2 | 2 | 2 | 2 | … | 7 | 7 | 7 | 7 |
| t | 0 | 1 | 2 | 3 | 0 | 1 | 2 | 3 | 0 | 1 | 2 | 3 | … | 0 | 1 | 2 | 3 |
8个warp × 每warp8组 = 64组,每组4线程,总 64×4 = 256 线程
二、16×128 矩阵切成 4 个 tile
1 | CHUNK=16, D=128 |
tile 布局:
1 | 列 0 列 63 列 64 列 127 |
tile_idx 计算:
1 | tile_idx = (m_blk/8) * 2 + (n_blk/64) |
循环顺序: tile0 → tile1 → tile2 → tile3
每个线程每次循环取2个元素,4次循环共取 4×2 = 8 个元素
三、每线程负责的行和列
1 | int row = m_blk + ((warp_id + g) % 8); |
-
g同时决定行和列段: -
列段: g=0 → col[0…7], g=1 → col[8…15], …, g=7 → col[56…63]
-
行:
(warp_id + g) % 8做循环移位
单个warp (warp0) 的分配:
| 段0 col0-7 | 段1 col8-15 | 段2 col16-23 | 段3 col24-31 | 段4 col32-39 | 段5 col40-47 | 段6 col48-55 | 段7 col56-63 | |
|---|---|---|---|---|---|---|---|---|
| row0 | w0,g=0 | |||||||
| row1 | w0,g=1 | |||||||
| row2 | w0,g=2 | |||||||
| row3 | w0,g=3 | |||||||
| row4 | w0,g=4 | |||||||
| row5 | w0,g=5 | |||||||
| row6 | w0,g=6 | |||||||
| row7 | w0,g=7 |
warp0 走对角线!
warp1 (row = (1+g)%8) 偏移一格:
| 段0 col0-7 | 段1 col8-15 | 段2 col16-23 | 段3 col24-31 | 段4 col32-39 | 段5 col40-47 | 段6 col48-55 | 段7 col56-63 | |
|---|---|---|---|---|---|---|---|---|
| row1 | w1,g=0 | |||||||
| row2 | w1,g=1 | |||||||
| row3 | w1,g=2 | |||||||
| row4 | w1,g=3 | |||||||
| row5 | w1,g=4 | |||||||
| row6 | w1,g=5 | |||||||
| row7 | w1,g=6 | |||||||
| row0 | w1,g=7 |
四、8个warp叠加: 循环移位矩阵
全部8个warp叠加后,每个格子恰好被一个warp负责:
| 段0 | 段1 | 段2 | 段3 | 段4 | 段5 | 段6 | 段7 | |
|---|---|---|---|---|---|---|---|---|
| row0 | w0 | w7 | w6 | w5 | w4 | w3 | w2 | w1 |
| row1 | w1 | w0 | w7 | w6 | w5 | w4 | w3 | w2 |
| row2 | w2 | w1 | w0 | w7 | w6 | w5 | w4 | w3 |
| row3 | w3 | w2 | w1 | w0 | w7 | w6 | w5 | w4 |
| row4 | w4 | w3 | w2 | w1 | w0 | w7 | w6 | w5 |
| row5 | w5 | w4 | w3 | w2 | w1 | w0 | w7 | w6 |
| row6 | w6 | w5 | w4 | w3 | w2 | w1 | w0 | w7 |
| row7 | w7 | w6 | w5 | w4 | w3 | w2 | w1 | w0 |
循环移位矩阵! 保证同一时刻每行每列段只有一个warp在访问。
五、一个格子 (1行×8列) 内4线程的分工
每个格子 8 列,由同组的 4 个线程 (t=0,1,2,3) 瓜分,每线程 2 列:
1 | 格子: row=R, col_base=C (8个元素) |
六、local_tile 图解
第一步: local_tile(g_tile, (1,8), make_coord(row, col_tile))
1 | g_tile [16行][128列]: |
含义:
- 把 [16,128] 按 (1行, 8列) 粒度切分
- 行方向: 每1行一个tile,共16个,编号0-15
- 列方向: 每8列一个tile,共16个,编号0-15
make_coord(row, col_tile)= 取第row行、第col_tile个列段- 结果
tile_g: shape=(1,8),指向g_tile[row][col_tile×8 .. col_tile×8+7]
第二步: local_tile(tile_g, (1,2), make_coord(0, t))
1 | tile_g: [g[3][16], g[3][17], g[3][18], g[3][19], g[3][20], g[3][21], g[3][22], g[3][23]] |
含义:
- 把 (1,8) 按 (1,2) 切分
- 列方向: 8/2=4个tile,编号0-3
make_coord(0, t)= 取第0行tile、第t个列tile- 结果
s_g: shape=(1,2),指向 shared memory 中的2个连续元素
对于 1D 的 g_total [128]:
1 | vec8_1d = (8) |
七、为什么用循环移位: 避免 bank conflict
Shared memory 有 32 个 bank,每 bank 4字节宽。bf16 元素 2字节,所以每 bank 放 2 个 bf16。
如果 warp 内 32 线程访问同一行连续 128 列:
| 线程 | 访问列 | bank |
|---|---|---|
| 0 | col0,1 | bank 0 |
| 1 | col2,3 | bank 1 |
| … | … | … |
| 31 | col62,63 | bank 31 |
完美无冲突!
但如果多线程访问同一列段的不同行:
可能多线程撞同一 bank → 串行化 → 慢!
循环移位保证: 同一 warp 内的 8 组 (g=0…7) 访问不同行、不同列段 → 落入不同 bank → 无冲突
八、完整流程总图
1 | Phase A: shared memory 存着 q[16][128], k[16][128], g[16][128] (gate cumsum后) |
九、计算公式
1 | q_decayed[row][col] = q[row][col] × exp2(g[row][col]) × scale |
代码实现:
1 | float g = reg_g[tile_idx][v]; |
注意: reg_gt 存的是 exp2(g_total[col]),在 gate+cumsum 阶段已经算好。
所以 k_restored = k × exp2(-g) × exp2(g_total) = k × exp2(g_total - g)
- 标题: KDA源码剖析之FlashKDA(上)
- 作者: 鱿鱼圈
- 创建于 : 2026-06-05 23:50:00
- 更新于 : 2026-06-14 15:45:15
- 链接: https://yuyanqi.com/2026/06/05/KDA源码剖析之FlashKDA(上)/
- 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。