不可不知小技巧|H800 上基于 Triton-TLE 优化 FlashKDA,长文本推理性能提升超 40%
编者按:线性注意力具有更低的长序列计算复杂度,但理论优势并不会自动转化为 GPU 上的实际性能。即使执行同一套 KDA 算法,线程是否长时间等待、状态保存在什么位置,以及数据需要搬运多少次,都会直接影响最终延迟。本文通过对比 FlashKDA 与 Triton-TLE,展示如何让并行准备阶段减少同步空转,并让顺序递推阶段避免主状态的反复搬运。在本文汇总的 12 组 H800 配置中,这些优化带来了 1.380× 的几何平均加速,在两个重点长文本配置上的加速约为 1.4×。
在大模型推理场景里,长文本越长,标准全注意力需要处理的数据就越多,计算量和显存占用也随之增长。Kimi Linear 采用 KDA(Kimi Delta Attention)与全注意力层 3:1 的混合结构,在模型效果与推理效率之间取得平衡。KDA 不必保存全部历史信息,而是将其压缩到一份固定大小的状态中,并通过更细粒度的门控决定哪些信息保留、哪些信息遗忘。它的状态更新可以拆成“按通道衰减”和“一次小规模校正”,从而把多个 token 合并成一个 chunk 处理,为 GPU 并行计算创造条件。
FlashKDA 将这套计算拆成两个阶段:K1 并行准备各个 chunk 所需的数据,K2 再按时间顺序更新状态。我们用 Triton-TLE 沿用这一框架,但分别优化了两条路径:K1 通过减少线程等待,让 GPU 计算单元更持续地工作;K2 则把反复使用的主状态留在寄存器中,减少它在不同存储层级之间的搬运。在本文汇总的 12 组 H800 配置中,TLE 相对 FlashKDA 取得 1.380× 的几何平均加速,在两个重点长文本配置上的加速约为 1.4×。
<01>FlashKDA 是什么?
Kimi Delta Attention 可以看成一条带门控的递归状态链。每个 token 先用 key 从旧状态中读出预测,再用 value 与预测之间的误差更新状态,最后由 query 读取输出。它绕开了标准 attention 的完整T × T矩阵,代价是引入了串行依赖:后一个 token 的状态必须等前一个 token 计算完成。
FlashKDA 的做法是固定BT=16,把序列切成num_chunks = ceil(seq_len / 16)个 chunk。chunk 内的 16 个 token 被整体改写为向量扫描和小矩阵运算,不再逐 token 推进完整状态;不同 chunk 的局部算子可以并行构造,只有携带状态的 chunk 间递推必须保持时间顺序。忽略部分缩放和尾块处理后,完整结构可以写成:
BT = 16num_chunks = ceil_div(seq_len, BT)# K1:所有 (sequence, head, chunk) 相互独立,可以同时执行。parallel_for (sequence, head, chunk):token_range = chunk * BT : (chunk + 1) * BTq_c, k_c, g_c, beta_c = slice_inputs(token_range) # [16, D]# 以下 scan 和 16×16 矩阵运算在 chunk 内并行处理 16 个 token。g_prefix = prefix_scan(safe_gate(g_c))Aqk[chunk] = build_causal_qk_matrix(q_c, k_c, g_prefix)Akk_inv[chunk] = invert_strict_lower_k_matrix(k_c, g_prefix, beta_c)w[chunk], qg[chunk], kg[chunk], decay[chunk] =build_state_terms(q_c, k_c, g_prefix, beta_c)# K2:不同 (sequence, head) 状态链可以并行,同一条链内的 chunk 必须串行。parallel_for (sequence, head):state = initial_state_or_zero()for chunk in range(num_chunks): # sequentialtoken_range = chunk * BT : (chunk + 1) * BTv_c, beta_c = slice_value_and_beta(token_range)v_new = Akk_inv[chunk] @ (sigmoid(beta_c) * v_c- w[chunk] @ state)output[token_range] = qg[chunk] @ state + Aqk[chunk] @ v_newstate = state * decay[chunk] + transpose(kg[chunk]) @ v_new
因此,K1 的启动网格覆盖所有(序列, head, chunk),用充足的 chunk 数量提供并行度;K2 的启动网格只覆盖(序列, head),每个 CTA 持有一条状态链并依次消费 K1 生成的工作区。这里的分块并未消除状态依赖,而是把依赖从逐 token 降低为每 16 个 token 一次,并把依赖之外的计算提前并行完成。
代码用H表示 Q/K head 数、HV表示 value head 数。这里讨论两种后端共同支持的HV == H路径,所以下文统一写作H。
Chunk KDA 两阶段结构
<02>原始设计:FlashKDA 的两阶段拆分与优化
FlashKDA 首先解决了两类工作不适合共享同一资源配置的问题。早期原型把 chunk 内计算和状态递推融合在一个内核中,前者的高并行度被后者的串行依赖限制,导致大量 streaming multiprocessor (SM)缺少可调度的工作。拆成 K1/K2 后,两个阶段可以分别选择启动网格、线程数和数据驻留策略;根据 FlashKDA v1 深入解析(https://github.com/MoonshotAI/FlashKDA/blob/d2ff19a/docs/20260420-flashkda-v1-deep-dive.md),这次拆分带来了至少 15% 的端到端收益。
两个阶段都以 BT=16 为基本计算粒度。这个尺寸既控制了安全门控变换(safe gate)在 BF16 下的数值范围,也使 chunk 内的L成为16 × 16严格下三角矩阵。由于L^16=0,单位下三角矩阵(I+L)的逆可通过有限 Neumann 展开计算;同时,16 × 16的矩阵尺寸也与 Tensor Core 计算块自然对齐。
K1:以 chunk 为单位并行准备
K1 面向大量相互独立的 chunk,优化目标是降低单个 Cooperative Thread Array(CTA )的资源占用,让更多 CTA 并发执行。每个 CTA 负责一个(序列, head, chunk),用 Tensor Memory Accelerator(TMA) 一次性将 q、k、g、beta 和门控偏置装入 shared memory,再依次完成归一化、门控前缀和、衰减变换、小矩阵乘和 Neumann 展开。
为贴近源码,下面沿用k_decayed/q_decayed/k_restored/g_total/INV/Mqk这些工作区名称;其中INV和Mqk分别对应前文简式中的Akk_inv和Aqk,其余张量共同编码 K2 读取旧状态和更新状态所需的变换。省略具体布局和尾块处理后,执行流程如下:
// 一个 8-warp CTA 处理一个 chunk。
flashkda_k1(chunk):TMA.load(q, k, g, beta, dt_bias -> shared_memory)cta_barrier()parallel_all_warps: normalize_qk_in_place()cta_barrier()parallel_threads: g_prefix, g_total = safe_gate_and_prefix_sum(g, dt_bias)cta_barrier()qk_gprefix_regs = load_from(smem.q, smem.k, smem.g_prefix)cta_barrier()// 两组数据生命周期不重叠,复用同一片 shared memory。write_qk_variants_to(reused_smem, qk_gprefix_regs, g_total)cta_barrier()warp_0: L = mma(k_decayed, k_inv)warp_1: Mqk = mma(q_decayed, k_inv)cta_barrier()apply_causal_masks_and_beta(L, Mqk, beta)cta_barrier()warp_0: INV = neumann_inverse(L)cta_barrier()TMA.store(k_decayed, q_decayed, k_restored,g_total, INV, Mqk -> workspace)
这里的 union 是资源控制的关键:q/k/g 完成衰减变换后不再使用,原来的存储空间随即交给k_decayed/q_decayed/k_inv/L/INV/Mqk等中间量,仅这一处生命周期复用就节省约 14 KB shared memory。K1 固定使用包含 8 个 warp 的 CTA;配合__launch_bounds__(256, 8),寄存器用量被压到每线程 32 个。使用 FP16 执行 Neumann 展开求逆、以 2 为底的ex2.approx和基于tanh.approx的 sigmoid,则进一步减少了 chunk 内小型计算的开销。
K2:用流水线顺序推进状态
K2 的约束完全不同:每个 CTA 对应一个(序列, head),必须沿 chunk 顺序推进同一份状态。FlashKDA 将 CTA 划分为 1 个加载 warp、4 个 MMA warp 和 1 个写回 warp,并设置 3 级输入流水线和 2 级输出流水线。
加载 warp 用 TMA 将 K1 工作区和当前 value 搬入 shared memory;MMA warp 等待当前输入槽位就绪,完成递推后提交输出槽位;写回 warp 再用 TMA 将结果写回全局内存。省略布局和边界处理后,三个角色的数据流如下:
// 一个 6-warp CTA 处理一条状态链。
shared BF16 state[D, D]input_pipe = pipeline(stages=3)output_pipe = pipeline(stages=2)state = load_initial_state_or_zero().to(BF16)load_warp:for chunk in chunks:slot = input_pipe.acquire(chunk)TMA.load(v, beta, k_decayed, q_decayed, k_restored,g_total, INV, Mqk -> slot)input_pipe.commit(chunk)mma_warps:for chunk in chunks:tile = input_pipe.wait(chunk)out_slot = output_pipe.acquire(chunk)prediction = k_decayed @ stateU = INV @ (sigmoid(beta) * (v - prediction))output = q_decayed @ state + Mqk @ Ustate = BF16(state * g_total + transpose(k_restored) @ U)store(out_slot, output)output_pipe.commit(chunk)input_pipe.release(chunk)store_warp:for chunk in chunks:out_slot = output_pipe.wait(chunk)TMA.store(out_slot -> global_output)output_pipe.release(chunk)
完整状态以 BF16 保存在 shared memory 中,因此 4 个 MMA warp 每处理一个 chunk 都从中读取旧状态,并将更新后的状态写回原处;如果接口要求 FP32 初始或最终状态,则只在循环前后进行格式转换。中间量 U 则用 MOVM_T 在寄存器内转换 MMA 分片布局,避免仅为布局转换而往返 shared memory。输入加载、状态递推和输出写回由两条流水线并行推进,但状态本身仍严格沿 chunk 顺序更新。
综上,FlashKDA 形成两种差异化资源配置策略:K1 依靠较小的单 CTA 资源占用扩大并发,K2 则以共享内存中的状态控制递推阶段的资源占用。这套设计兼顾了占用率与不同配置下的扩展性。Triton-TLE 沿用两阶段结构,但针对 H800 的目标负载,在关键路径上采用了不同的资源取舍。
<03>瓶颈剖析:从 FlashKDA 到 Triton-TLE 的设计动机
在 H800 的目标负载上,FlashKDA 的两个阶段呈现出不同的资源特征。以定长 N=1, H=96, T=1024 为例,K1 已有接近满载的占用率,却没有转化为同等水平的指令发射;K2 的启动网格小于 SM 数量,继续压低单 CTA 的资源占用也无法增加可并行的状态链。Triton-TLE 因而没有为两个阶段套用同一种优化策略。
这两个阶段的优化思路不一样:K1 依靠大量独立 chunk 提供并行度,需要让已有 warp 更持续地发射有效工作;Triton-TLE 用显式共享内存张量和异步加载组织跨阶段的数据复用,再配合张量级矩阵运算与资源自动调优选择更紧凑的 CTA。K2 无法并行展开相邻 chunk,需要把更多片上资源集中到单条状态链上;这里输入加载和输出写回仍可与 MMA 重叠,因此 Triton-TLE 用 TensorDescriptor、tle.gpu.copy、流水线和 warp 特化(warp specialization)把外围搬运组织在寄存器递推周围。
需要强调的是,这不是一次把整个 K2 迁移到 WGMMA 的优化。两种实现都使用 TMA,在 Hopper 上,Triton-TLE 只有kg^T @ v_new这一步状态更新使用 WGMMA。设计核心始终围绕 K1 的有效发射率和 K2 的状态驻留方式。
<04>为什么选择 Triton-TLE:将数据流与执行角色显式编码
Triton 本身已经能够用tl.dot表达张量计算,用triton.autotune搜索内核配置,也能通过 TensorDescriptor 使用 TMA。Triton-TLE 的价值并不是提供 Triton 无法访问的硬件指令,而是在这些能力之上,把共享内存对象、异步搬运、缓冲槽位、执行角色及其资源分配组织成可以组合的一等编程抽象。因此,这里所说的“相对于 Triton 的独有能力”,指的是 Triton-TLE 补充的程序结构,而不是对 TMA 或 WGMMA 等硬件能力的独占。对于 KDA 这样同时包含高并行准备阶段和串行状态链的算子,这些抽象使数据放在哪里、由谁搬运、何时可被下一角色消费,都能直接体现在程序结构中。
本文实现用到的主要能力如下:
这些能力在两个阶段承担的职责并不相同。K1 借助显式共享内存张量和异步加载组织数据复用,再与 Triton 的tl.dot和自动调优配合,选择更适合 chunk 内计算的 CTA 配置。K2 使用tle.pipe与tle.gpu.warp_specialize,pipe 把w/v/qg/kg/Aqk/Akk/gk等多组输入作为一个有明确生命周期的数据槽位传递,warp 专门化则把外围搬运与顺序递推分给不同角色。
FP32 主状态跨 chunk 驻留在寄存器中,本身是一项数据驻留决策,并非单独的 Triton-TLE原语。Triton-TLE 的作用是让整个递推循环留在 MMA 消费者中,同时用结构化的数据流把加载和写回围绕它组织起来。下面分别说明这些抽象如何映射到 K1 和 K2。
<05>Triton-TLE 实现:把资源用在关键路径上
K1:从高占用率转向高有效发射率
K1 的任务是把每个 chunk 的原始q/k/g/beta转换为一组 K2 可以直接使用的小矩阵和向量。不同 chunk 之间没有状态依赖,所有(序列, head, chunk)独立并行。
K1:chunk 内算子构造
简化后的算法如下:
for each chunk in parallel:
q, k = l2_normalize(q), l2_normalize(k)g_prefix = prefix_scan(safe_gate(g))Aqk = causal_lower((q * exp(g_prefix)) @ (k * exp(-g_prefix))^T)L = strict_lower((k * exp(g_prefix)) @ (k * exp(-g_prefix))^T)Akk_inv = neumann_inverse(L * sigmoid(beta))emit w, qg, kg, g_last, Aqk, Akk_inv
其中,Aqk描述 chunk 内 value 对输出的直接贡献,Akk_inv用来修正 value;w/qg/kg/g_last则把旧状态的读取、更新和衰减整理成 K2 需要的形式。K1 完成后,K2 不再处理门控扫描或三角求逆,只需执行形状规则的矩阵乘和状态递推。
FlashKDA K1 的 8 个 warp 并非在每个阶段都有等量工作。归一化和衰减阶段可以利用较多线程,L/Mqk的 MMA 和 Neumann 逆矩阵计算却只激活少数 warp;阶段之间的全 CTA 屏障又让其余常驻 warp 等待。结果是占用率很高,但每个周期的有效指令发射仍然有限。
K1 不需要 K2 那样的生产者—消费者流水线。Triton-TLE 在这里的直接作用,是显式分配原始 BF16q[16,128]、原始 BF16k[16,128]和 FP32 门控前缀和三个共享内存张量,并用异步加载填充它们,供后续阶段重复读取。归一化系数保存在寄存器中,Aqk/Akk和逆矩阵计算则用张量级tl.dot表达。
省略边界处理和具体布局后,实现思路可以写成下面的伪代码;它只描述数据驻留和计算组织,不对应完整的 API 签名:
@triton.autotune(configs=k1_resource_configs, ...)
def k1(...):q_buf = tle.gpu.alloc([BT, K], dtype=bf16, scope=tle.gpu.smem)k_buf = tle.gpu.alloc([BT, K], dtype=bf16, scope=tle.gpu.smem)gc_buf = tle.gpu.alloc([BT, K], dtype=fp32, scope=tle.gpu.smem)q_ptr = tle.gpu.local_ptr(q_buf, tile_indices)k_ptr = tle.gpu.local_ptr(k_buf, tile_indices)gc_ptr = tle.gpu.local_ptr(gc_buf, tile_indices)q = tle.load(q_block, is_async=True)k = tle.load(k_block, is_async=True)g = tle.load(g_block, is_async=True)tl.store(q_ptr, q)tl.store(k_ptr, k)tl.store(gc_ptr, prefix_scan(safe_gate(g)))q_rstd, k_rstd = l2_rstd(q), l2_rstd(k)Aqk = tl.dot(q_with_decay, transposed_k_with_decay)Akk = tl.dot(k_with_decay, transposed_k_with_decay)Akk_inv = neumann_inverse_with_tl_dot(Akk)# 再次读取共享内存中的 q/k/g,构造 K2 所需的工作区。emit_workspace(q_ptr, k_ptr, gc_ptr, q_rstd, k_rstd)store_outputs(Aqk, Akk_inv)
这里tle.gpu.alloc/local_ptr和异步tle.load是 Triton-TLE 提供的数据驻留与搬运原语;tl.dot和自动调优来自 Triton。二者结合后,K1 不必沿用 FlashKDA 固定的 8-warp 阶段划分,而可以在包含 2/4/8 个 warp 的 CTA 配置间选择;这里讨论的负载最终采用包含 4 个 warp 的 CTA。
这种选择并没有减少指令数:代表配置中,Triton-TLE K1 执行38.51M条指令,高于 FlashKDA 的29.99M。它的收益来自更少的全 CTA 同步和角色空转,使issue active从39.45%提升到63.69%,barrier stall/issue从20.21降到0.89。因此,包含 4 个 warp 的 CTA 虽然使用更多寄存器、理论占用率也更低,Tensor Core 工作却更连续。
Aqk/Akk采用[batch, head, chunk, 16, 16]布局,使 K1 写出的每个16 × 16矩阵与 K2 读入的对应矩阵都在内存中形成连续的数据块。
K2:让 FP32 状态留在寄存器里
K2 中每个(序列, value head)对应一条状态链。输入加载和输出写回可以通过流水线重叠,但 chunkc+1必须等 chunkc产生新状态——循环内部的递推路径才是真正要缩短的目标。
K2:state recurrence 与流水线
Triton-TLE K2 保留了生产者—消费者流水线,角色配置为 4 个加载 warp、4 个 MMA warp 和 1 个写回 warp,输入和输出各使用四级流水线。最关键的变化是:完整的 FP32 主状态在整个 chunk 循环中始终留在 MMA 消费者的寄存器上下文里。
这里需要区分逻辑角色与实际启动配置。加载/MMA/写回的逻辑划分是4/4/1,但硬件以分区粒度分配 warp,最终 CTA 包含 384 个线程,即 12 个物理 warp。代码中外层num_warps=4配置默认的加载分区,warp_specialize的[4, 1]分别指定 MMA 和写回工作分区。
实现层面,Triton-TLE 不只是换一种方式调用tl.dot。TensorDescriptor 和tle.gpu.copy负责搬运数据块,两个容量为 4 的tle.pipe管理输入与输出,tle.gpu.warp_specialize再把加载、MMA 和写回工作分给不同 warp。省略参数和边界处理后,这套生产者—消费者模型可以概括为下面的角色级伪代码:
load_pipe = tle.pipe(..., capacity=4)
store_pipe = tle.pipe(..., capacity=4)def load_producer():for chunk in chunks:slot = load_pipe.acquire(chunk)tle.gpu.copy(input_descriptors, slot)load_pipe.commit(chunk)def mma_consumer():state = load_initial_state().to(tl.float32)for chunk in chunks:tile = load_pipe.wait(chunk)output, state = recurrence(tile, state)out_slot = store_pipe.acquire(chunk)store(out_slot, output)store_pipe.commit(chunk)load_pipe.release(chunk)def store_consumer():for chunk in chunks:output = store_pipe.wait(chunk)tle.gpu.copy(output, output_descriptor)store_pipe.release(chunk)tle.gpu.warp_specialize(load=(load_producer, 4),mma=(mma_consumer, 4),store=(store_consumer, 1),)
加载生产者最多提前准备 4 个 chunk,写回消费者也能独立写回已完成的输出。它们可以和 MMA 发生重叠,但不会改变状态递推的依赖关系:只有持有 FP32 主状态的 MMA 消费者必须严格沿 chunk 顺序推进。Triton-TLE 的作用,是把可并行的搬运和写回组织在这条串行路径周围。
在这个执行框架内,MMA 消费者的核心递推如下:
# b_h is the FP32 master state and lives across the whole loop.
b_h = load_initial_state().to(tl.float32)for chunk in chunks:w, qg, kg, Akk_inv, Aqk, g_last, v_beta = load_pipe.wait(chunk)# Stage a BF16 operand in shared memory for HMMA.b_h_bf = b_h.to(tl.bfloat16)# Correct the value with the old state.kh = tl.dot(w, b_h_bf).to(tl.float32)v_new = tl.dot(Akk_inv, (v_beta - kh).to(tl.bfloat16)).to(tl.float32)# Read the old state and add the chunk-local output.output = scale * tl.dot(qg, b_h_bf)output += tl.dot(Aqk, v_new.to(tl.bfloat16))# Keep the updated master state in FP32 registers.b_h = b_h * exp2(g_last)[:, None]b_h += tl.dot(tl.trans(kg), v_new.to(tl.bfloat16)).to(tl.float32)
这份 FP32 主状态始终跨 chunk 驻留在 MMA 消费者的寄存器里。为了供 HMMA 读取,每个 chunk 会生成一份128 × 128的 BF16 操作数,暂存在共享内存中参与w @ h和qg @ h;kg^T @ v_new则由 WGMMA 以 BF16 输入、FP32 累加的方式直接更新寄存器中的主状态。Triton-TLE 主要避免了将更新后的主状态反复写回共享内存再重新读出。
FlashKDA 与 Triton-TLE 的数据驻留差异
相应地,K2 每线程使用 168 个寄存器,每个 SM 最多只能驻留一个 CTA。不过在 H800 上定长N=1, H=64/96的主要负载中,这并不会增加调度波次:H800 有 132 个 SM,64 或 96 个 K2 CTA 都能在一个调度波次内启动。此时,更多寄存器是避免主状态反复往返共享内存的代价,不能脱离上下文直接判定为性能损失。K2 的实际启动网格为N × H;变长或批处理场景还要考虑序列数N。
<06>性能结果与适用边界
性能对比主要在 132 个 SM 的 NVIDIA H800 上跑,输入 BF16,K=V=128、BT=16。端到端延迟由 CUDA 事件测量,NCU 重放耗时(replay duration)只用于区分 K1/K2 的阶段贡献。完整测试可见文末源码仓库。
在汇总的 12 个 H800 配置中,Triton-TLE 均快于 FlashKDA,几何平均加速为1.380×,范围为1.272×–1.439×。当T=8192时,重点关注的H=64/96负载分别达到1.419×和1.401×:
H800 上 Triton-TLE 相对 FlashKDA 的加速
H=96,T=1024的阶段数据进一步说明,两个内核的收益来源并不相同:
在这个配置上,K1 自身加速1.264×,减少22.240 us;K2 自身加速1.474×,减少47.424 us。按两段 NCU 重放耗时的减少量计算,K1 和 K2 分别贡献约32%和68%。这一比例仅用于说明 K1 和 K2 各自对阶段耗时下降的贡献,不应视为对端到端加速的精确拆分。
H20 的 78 个 SM 则揭示了这套资源选择的边界。对于 Triton-TLE 生产版本后端,在定长N=1、逐步增加 head 数的测试中,H=64时仍有1.136×加速;到H=96,每个 SM 只能驻留一个 CTA 的 Triton-TLE K2 至少需要第二个调度波次,加速随之降到1.035×。在汇总的 12 个配置中,Triton-TLE 有 11 个领先,几何平均加速为1.050×。
H20 上的加速与调度边界
<07>结语
回顾一下整条优化路径,FlashKDA 已经完成了最关键的一步:用 BT=16 在数值稳定性和计算效率之间找到了平衡点,并把 token 并行的准备阶段和 head 并行的递推阶段拆开了。Triton-TLE 沿 着这套结构继续优化,K1 用显式 shared memory 张量、异步加载和张量级计算撑起更紧凑的 CTA,让已有的 warp 能持续发有效指令;K2 则把 FP32 主状态跨 chunk 留在寄存器里,用生产者-消费者流水线把加载和写回安排在串行递推的周围。
这次优化最值得记住的经验是:寄存器、占用率、屏障等待这些指标,都不能脱离具体负载去解读。在 H800 的重点场景里拿到了约 1.4× 的端到端加速,H=96,T=1024 的阶段数据也说明 K1 和 K2 都有贡献——K1 靠减少同步空转改善了 chunk 内计算的有效发射,K2 靠寄存器驻留主状态减少了每轮递推的指令和 shared memory 往返,生产者-消费者模型负责把外围搬运重叠起来。对定长 N=1, H=64/96,K2 的资源选择没有增加调度波次;但在只有 78 个 SM 的 H20 上,同一选择从 H=96 开始就会引入额外波次。性能判断终究要回到目标配置的关键路径,结合资源配置、调度波次和消融结果来解释性能剖析数据。
本次优化源码地址:
https://github.com/flagos-ai/FlagGems-vllm/blob/main/src/flaggems_vllm/ops/FLA/chunk_kda.py
FlashKDA v1 深度解析:
https://github.com/MoonshotAI/FlashKDA/blob/d2ff19a/docs/20260420-flashkda-v1-deep-dive.md