资讯动态

PyPTO-Gym 的 chunk_kda 算子实战:chunked gated delta-rule 线性注意力的 NPU 前向实现

发布时间:2026/9/19 13:51:45 来源:尧图企业网站定制
PyPTO-Gym 的 chunk_kda 算子实战chunked gated delta-rule 线性注意力的 NPU 前向实现【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读chunk_kda是 CANN PyPTO-Gym 仓库中面向 Kimi Delta AttentionKDA带每维遗忘门g的 delta-rule 线性注意力的前向forward-only融合算子位于 src/pypto_gym/ops/pypto_tensor/ling_3_0_flash/chunk_kda/。它以chunk_size128切分序列在 chunk 内构造下三角矩阵A并用 8×8 分块前向代入完成 delta 修正在 chunk 间以 fp32 状态矩阵S递推从而把 O(T²) 的注意力退化为线性复杂度。读完本文你将掌握该算子的完整 ABI输入输出形状与 dtype 契约、数值稳定性设计原理、固定长度BSND与 varlenTND两种调用方式以及仓库中配套的 golden 验证与测试方法。算子定位与算法背景chunk_kda属于ling_3_0_flash子目录下两个 KDA 融合算子之一另一个是逐 token 递推的 fused_recurrent_kda两者共享D K V 128的公共形状约束见 ling_3_0_flash/README.md。与标准 softmax 注意力不同delta-rule 线性注意力通过每维遗忘门glog 域≤0控制状态衰减并用beta权重对 key/value 做 delta 修正更新规则为S S * exp(g) # 门控衰减 S outer(beta * k, v - (k · S)) # delta 更新 o q · S # 注意力输出朴素实现按 token 逐次递推即 recurrent 路径而chunk_kda采用分块chunked策略一次处理 128 个 token块内通过下三角矩阵求逆一次性完成所有 delta 修正块间仍以状态S递推兼顾精度与 NPU 上的矩阵乘利用率。产品支持情况按 chunk_kda/README.md 的声明Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持实现代码中通过pypto.frontend.jit的runtime_optionsrun_modepypto.RunMode.NPU强制 NPU 运行模式且依赖torch_npu完成 NPU 设备初始化见 chunk_kda_impl.py。算子签名与 I/O 规格chunk_kda的公开入口是主机侧 wrapper 函数签名如下与 README 及实现 chunk_kda_impl.py 一致chunk_kda_wrapper(q, k, v, g, beta, scaleNone, initial_stateNone, output_final_stateFalse, use_qk_l2norm_in_kernelFalse, cu_seqlensNone, **kwargs) - (o, S)各输入输出的形状与 dtype 契约I/O变量ShapeDtype说明入q, k, g[B, T, H, K]bf16query / key / 每维遗忘门log 域入v[B, T, H, V]bf16value入beta[B, T, H]fp32delta 更新权重入scalescalarfloatNone→K**-0.5入initial_state[N, H, V, K] / Nonefp32初始状态 SNone时清零出o[B, T, H, V]bf16注意力输出出S[N, H, V, K] / Nonefp32output_final_stateTrue时返回需要特别说明两点 ABI 细节均可在源码中得到印证g与beta必须是 fp32 且由调用方保证 dtype实现中的_validate_inputs会断言q/k/v为torch.bfloat16、g/beta为torch.float32chunk_kda_impl.pywrapper 不再做 dtype 转换。initial_state/S采用上游 ABI 布局 [N, H, V, K]即 S 的转置而 kernel 内部计算使用 [K, V][V,K] ↔ [K,V]的转置在 kernel 内部完成主机侧不做搬运chunk_kda_impl.py。支持的 dtype 与形状约束dtypebf16 输入内部全程 fp32输出o为 bf16状态S保持 fp32。T % 128 0当cu_seqlensNone时varlen 模式支持非对齐尾部 chunk尾部不足 128 的行在 kernel 内通过pypto.fillpad以 0 填充无需主机侧 pad。K V 128P0 阶段硬约束_validate_inputs末尾断言K 128 and V 128见 chunk_kda_impl.py。容差bf16 输出o的 rtol/atol 为 1e-2fp32 状态S的容差为 1e-3。算法原理chunk 内求逆 chunk 间状态递推从 golden 参考实现 chunk_kda_golden.py 的模块划分M1–M4可以清晰还原整个计算流程。golden 的语义与theta-flash-linear-attention/fla/ops/kda/naive.py::naive_chunk_kda及kda.py::chunk_kda_fwd对齐见该文件头部注释。M1重排 缩放 门控 cumsum将[B,T,H,*]按BT128切分为NT T//128个 chunk 并重排为[B,H,NT,BT,*]q乘上scaleg在 chunk 内做前缀和得到g_cum。golden 中不使用torch.cumsum而是用下三角含对角矩阵左乘实现前缀和g_cum tril_incl gchunk_kda_golden.pykernel 侧同样用pypto.matmul(tril_incl, gc, pypto.DT_FP32)完成chunk_kda_impl.py。M2chunk 内 A 矩阵与 (IA) 求逆chunk 内衰减矩阵按A_full Σ_d k_c·k_i·exp(g_c − g_i)构造严格下三角再乘行beta得到A进而求解(IA)的逆golden 使用torch.linalg.solve_triangular一次性求解chunk_kda_golden.pykernel 侧则手工实现了8×8 分块分层求逆_inverse88 个 16×16 对角块各自前向代入求逆再按 16→32→64→128 的层次用_inv_merge合并inv22·A21·inv11的 Schur 补形式见 chunk_kda_impl.py。M3w / u 投影由A_inv与k、v计算 delta 修正后的投影w A_inv (exp(g_cum)·k)、u A_inv vchunk_kda_golden.py。M4跨 chunk 状态递推对每个 chunki先算v_i u_i − w_i S减掉旧状态贡献的 delta 修正输出o qg S A2 v_i再按S S·exp(g_last) (exp(g_last−g)·k)^T v_i递推下一 chunkchunk_kda_golden.py。kernel 侧的对应逻辑在_chunk_compute中chunk_kda_impl.py其中u与wS均以 fp32 累加qgS与a2vi两项累加得到最终oc。源码结构单 JIT kernel 主机 wrapper整个算子是一个模块级pypto.frontend.jitkernelchunk_kda_varlen_kernel加一个主机侧 wrapperchunk_kda_wrapper的组合文件为 chunk_kda_impl.py。JIT 编译配置kernel 的runtime_options与pass_options是典型的 NPU 融合算子调优参数chunk_kda_impl.pypypto.frontend.jit( runtime_options{ run_mode: pypto.RunMode.NPU, stitch_function_max_num: 256, device_sched_mode: 1, launch_sched_aicpu_num: 3, }, pass_options{ vec_nbuffer_setting: {-2: 1, -1: 32}, cube_l1_reuse_setting: {-1: 32}, cube_nbuffer_setting: {-1: 8}, }, )其中vec_nbuffer_setting控制向量单元多缓冲-2 保留 1 份、默认类 32 份cube_l1_reuse_setting与cube_nbuffer_setting控制 Cube 单元的 L1 复用与多缓冲服务于长序列下 chunk 循环的流水。kernel 内部结构kernel 的输入张量依次为q, k, v, g, beta, cu, tril, trils, eyestk, state0, o_out, s_out末尾跟随两个非张量的 trace 期参数use_qk_l2norm_in_kernelbool与scalefloat——它们共同构成编译缓存的关键字chunk_kda_impl.py。H头数与N序列数全部从张量维度在 kernel 内推导因此不同 (H,N) 复用同一份编译产物。计算主循环为两层pypto.loop外层LHhead 循环×LNsequence 循环内层LNT是以_BT128为步长的动态边界、静态步数 chunk 循环unroll_list[4]展开 4 份以提升指令级并行。每个 chunk 内用 rank-4 strided view 切出头、序列对应的 128 行切片并用valid_shape标记尾部 chunk 的有效行数chunk_kda_impl.py。尾部 chunk 的零填充这是 varlen 支持的关键设计pypto.is_loop_end(s_idx)分支内对最后一个可能不满 128 行chunk 用pypto.fillpad(..., constant, 0.0)将 pad 行清零——特别注意g的 pad 为 0这保证g_cum在 pad 区保持平坦、glast取值正确beta的 pad 同样为 0使 pad 行在 delta 修正中惰性化chunk_kda_impl.py。非尾部 chunk 则跳过 fillpad 直接计算。主机侧 wrapper 与输入校验chunk_kda_wrapper只做布局与分配随后发起一次 JIT 调用chunk_kda_impl.pycu_seqlensNone固定长度断言T % 128 0合成N B段、每段长T的cu [0, T, 2T, ..., B*T]int32设备张量且按(B, T, device)缓存复用并将[B,T,H,*]reshape 打包为[1, B*T, H, *]cu_seqlens给定varlen要求B 1cu原样透传保持设备端 int32 张量、kernel 内直接索引N len(cu) - 1initial_state按上游 ABI [N,H,V,K] 直接透传仅做contiguous()None时在设备端生成 fp32 零张量作为 seed三个 kernel 常量tril、trils、eyestk由_build_masks按设备缓存生成只依赖_BT/_MIN而不依赖输入数据chunk_kda_impl.py。_validate_inputs只做自省式introspection-only校验全部用.dim()/.shape/.dtype/.device/.numel()不做任何张量运算也不把cu_seqlens拷回主机它的值由 kernel 内索引消费校验只关心形状与 dtype见 chunk_kda_impl.py。使用示例以下是一个固定长度调用的最小示例输入构造方式对齐仓库测试 test_inputs.py 的稳定输入约束import os import torch import torch_npu # 必须导入以初始化 NPU 设备 from ling_3_0_flash.chunk_kda.chunk_kda_impl import chunk_kda_wrapper torch.npu.set_device(int(os.environ.get(TILE_FWK_DEVICE_ID, 0))) B, T, H, K 1, 4096, 16, 128 # T % 128 0 dev torch.device(npu:0) # 稳定输入q/k/v ~ randn*0.1g 取 logsigmoid(randn)log 域 ≤0beta 取 sigmoid(randn) g0 torch.randn(B, T, H, K, devicedev) q torch.randn(B, T, H, K, devicedev, dtypetorch.bfloat16) * 0.1 k torch.randn(B, T, H, K, devicedev, dtypetorch.bfloat16) * 0.1 v torch.randn(B, T, H, K, devicedev, dtypetorch.bfloat16) * 0.1 g torch.nn.functional.logsigmoid(g0).float() # fp32 beta torch.sigmoid(torch.randn(B, T, H, devicedev)).float() # fp32 o, S chunk_kda_wrapper(q, k, v, g, beta, scaleNone, # None - K**-0.5 initial_stateNone, output_final_stateTrue) # True 时返回 fp32 S print(o.shape, o.dtype) # [1, 4096, 16, 128] torch.bfloat16 print(S.shape, S.dtype) # [1, 16, 128, 128] torch.float32[N,H,V,K] 上游 ABIvarlenTND用法令B1将各序列打包拼接并传入cu_seqlens作为设备端 int32 张量segs [64, 200, 320] # 各序列长度允许非 128 对齐 cu torch.tensor([0, 64, 264, 584], devicedev, dtypetorch.int32) T 584 o, S chunk_kda_wrapper(q, k, v, g, beta, cu_seqlenscu, initial_stateNone, output_final_stateTrue)数值稳定性fp32 求逆与输入缩放README 明确列出两条已知约束二者在源码中均有硬性保证输入稳定性内部精度测试必须使用缩放后的稳定输入q/k/v ~ *0.1、g logsigmoid(randn) ≤ 0、beta sigmoid。test_inputs.py 将其标注为“HARD CONSTRAINT”未缩放的 raw randn 输入会使前向代入求逆出现 NaN近 1 的 gate 导致衰减指数动态范围过大。fp32 求逆硬约束(IA)与 8×8 分块前向代入求逆全程必须在 fp32 下进行bf16 会直接 NaNS累加器禁止窄于 fp32。实现中_intra_chunk_a的两个 band matmul 显式指定pypto.DT_FP32注释注明“极端衰减动态范围在 bf16 下会舍入”chunk_kda_impl.py。此外还有两个支撑稳定性的实现细节局部 pivot 分块_intra_chunk_a把 128×128 因果矩阵按_NC个 64 行块切分每块取块首行gp gcum[r0]作为局部 pivot行因子exp(gcum−gp)指数恒 ≤0列因子exp(gp−gcum)被_DECAY_CAP 80.0截断避免 fp32 exp 溢出约 88.7见 chunk_kda_impl.py。掩码常量折叠trils直接预置为-1严格下三角取负省去 kernel 内的取负运算chunk_kda_impl.py。测试与验证体系测试位于 tests/ops/ling_3_0_flash/chunk_kda/由四个文件组成chunk_kda_golden.pyStage-2 golden 参考实现按 M1–M4 模块划分支持固定长度cu_seqlensNone与 varlenTND两条路径并实现了use_qk_l2norm_in_kernel对 q/k 在特征维 K 上做逐 token L2 归一化eps1e-6。文件自带_validate()自检固定长度路径与naive_chunk_kda在 fp321e-3与 bf161e-2下对拍varlen 路径与逐序列naive_recurrent_kda对拍。test_chunk_kda.pyE2E 精度门禁按叶子输出逐项对比leaf0o容差 1e-2leaf1S容差 1e-3用例覆盖固定长度L0B1/T128/H2、T1B1/T512/H32、P0B1/T4096/H16、L2B4/T8192/H32、U0B1/T384/H2、U1B1/T1920/H32以及长序列压力用例 L3B1/T131072/H32默认运行列表之外需显式调用varlenV0单段与固定长度逐位相等、V1/V2/V3多段 随机初始状态 非对齐尾部、V4单段 T4117 非 %64 的 kernel 内尾 pad及 T37 的纯尾部 chunk 用例特性开关use_qk_l2norm_in_kernelTrue与 False 输出必须不同差值 1e-4证明该路径真实生效确定性同一输入两次运行逐位相等§S4NaN 顺序门禁按历史触发顺序重跑要求所有 launch 0 NaN / 0 Inf回归旧的陈旧 UB 残留 NaN 缺陷。test_inputs.py对抗性输入生成器直接在 NPU 上构造稳定输入另提供cancellation_stress模式针对 M3 中v u − wS的减法消去做 64 次随机缩放搜索。detailed_tensor_compare.py逐元素对比工具输出超差元素数量、比例、max/mean/std 差值与离群点明细。运行方式NPU 环境pytest tests/ops/ling_3_0_flash/chunk_kda/test_chunk_kda.py与其他 KDA 算子的关系chunk_kdaprefill/长序列前向与 fused_recurrent_kdadecode/逐 token 递推是ling_3_0_flash下互补的两个 KDA 算子前者以分块求逆换取并行度后者维护[H,V,K]的 fp32 状态矩阵做逐 token 递推并额外支持 inplace 状态ssm_state_indices、spec decodenum_accepted_tokens与 kernel 内 q/k L2 归一化。两者共享DKV128、Atlas A2/A3 系列支持的约束可分别服务于 KDA 注意力推理的 prefill 与 decode 阶段。如需了解其在完整模型中的接入方式可参考 ling_3_0_flash/README.md 中提到的 monkey-patch hook 模式。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价