- 人工智能
- 大模型
- 模型推理服务
- Ascend
- CANN
【免费下载链接】vllm-ascend
Community maintained hardware plugin for vLLM on Huawei Ascend
本篇技术指南围绕 vllm-ascend 开源仓库中的ChunkKdaFwd自定义算子展开,它实现了对齐 flash-linear-attention(fla-org)chunk_kda_fwd顶层语义的 Chunked KDA(Key-Decay Attention)前向计算,并在昇腾 Ascend 910B/910_93/950 系列 NPU 上以 L0(AscendC 单核算子)方式落地。读完本文,你将掌握该算子的 gate 数学公式、12 路返回值的完整契约、Python/aclnn 两种调用方式、L2 调度与四阶段内核流水(Gate/Prepare/Post-WU/FwdH/Finalize),以及 tiling key、重计算策略和验证矩阵等实现细节。
算子定位:ChunkKdaFwd 是什么
ChunkKdaFwd位于仓库 csrc/attention/chunk_kda_fwd,是 vllm-ascend 为昇腾 NPU 实现的 KDA 前向算子。它对齐不涉及 CP(Context Parallel)切分的 FLAchunk_kda_fwd顶层语义:公共接口既可以接收 raw gate,也可以接收已在 Python 层激活好的自然对数 gate;Gate、Prepare、PostWu、FwdH 和 Finalize 五个阶段均在一个物理ChunkKdaFwdL0 内核内完成,L2 层不再拼接或依次发射多个阶段 L0(A5 多 chunk 场景的特殊调度见下文“模板化方案与 tiling key”一节)。
从工程调用链看,vllm-ascend 在 vllm_ascend/ops/kda.py 的run_chunk_kda中通过torch.ops._C_ascend.chunk_kda_fwd真实使用该算子:以layout="BSND"、chunk_size=64(对应文件顶部的KDA_CHUNK_SIZE = 64)、state_v_first=True、use_gate_in_kernel=True运行,并配合safe_gate与lower_bound传入安全 gate 配置,输出(output, final_state)供上层复用 VK 状态。
本文涉及的形状符号沿用 KDA 约定:B为 batch,T为序列长度,N_c = T / chunk_size为 chunk 数,H为 q/k 的 head 数,H_v为 v 的 head 数(支持 GQA,H_v >= H且整除),K为 key 维度,V为 value 维度,N为变长序列条数。完整的 Shape 说明可在 API 文档 与 设计文档 中按本文后续各节对应查看。
Gate 公式与数学语义
令x = g + dt_bias,算子逐 token、逐 K 维计算自然对数衰减 gate。三种模式由属性use_gate_in_kernel与safe_gate组合决定:
use_gate_in_kernel = false: gate = g use_gate_in_kernel = true, safe_gate = false: gate = -exp(A_log) * softplus(x) use_gate_in_kernel = true, safe_gate = true: gate = lower_bound * sigmoid(exp(A_log) * x)随后在每个 chunk 内做 chunk-local 累计,再除以ln(2)转成 log2 域:
gk_i = cumsum(gate)_i / ln(2)因此后续exp2(gk)与自然指数 gate 严格绑定,不暴露额外的 gate scale——这是该算子数学语义的关键约束:gk 已经是“以 2 为底”的累计衰减量,任何外层都不再乘除额外系数。
几点值得注意:
safe_gate=true时,gate 通过 sigmoid/lower_bound 构造,保证数值落在稳定区间(lower_bound默认-5.0,取值范围[-5, 0)),避免exp溢出;use_gate_in_kernel=false时仍可搭配safe_gate=true,即外部已激活 gate 走后续稳定计算路径,safe 语义与 gate 生成位置相互独立;- 公共接口的语义优先级以稳定 Python / aclnn 接口为准(见“调用途径”一节)。
输入契约
算子输入如下表所示(layout只描述 q/k/v/g/beta 这些输入张量的布局;BSND/TND 由 L2 使用l0op::Transpose转为内部 BNSD/NTD 再进入内核):
| 名称 | 必选性 | Shape/Dtype | 说明 |
|---|---|---|---|
q/k | 必选 | 输入 layout 对应 Shape;FP16/BF16 | Query/Key |
v | 必选 | 输入 layout 对应 Shape;与 q 同 dtype | Value |
g | 必选 | 输入 layout 对应 K 维 Shape;FP32/BF16 | raw gate 或已激活自然对数 gate |
beta | 必选 | 去掉 g 的 K 维;FP32/BF16 | Delta 系数 |
A_log | 条件必选 | [H_v],FP32 | use_gate_in_kernel=true时必选 |
dt_bias | 可选 | [H_v*K],FP32 | gate bias |
initial_state | 可选 | [N,H_v,K,V]或[N,H_v,V,K],FP32 | 由state_v_first解释 |
cu_seqlens | 可选 | [N+1],INT64 | 变长序列边界 |
chunk_indices | 可选 | [2*N_c],INT64 | canonical chunk 顺序 |
从算子定义源码 op_host/chunk_kda_fwd_def.cpp 可以看到:q/k/v支持DT_FLOAT16/DT_BF16,g支持FP32/BF16,beta支持FP32/BF16,而a_log/dt_bias/initial_state固定为 FP32,cu_seqlens/chunk_indices固定为 INT64,所有输入均注册为FORMAT_ND并支持动态 shape(DynamicShapeSupportFlag(true))。
输出契约与 12 返回值语义
Python 层返回顺序固定为:
(attn_out, final_state, gk, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state)各返回值语义:
attn_out固定为 BSND/TND(即输出永远按“公开 layout”排布);final_state固定按序列排列,末两维服从state_v_first([N,H_v,K,V]或[N,H_v,V,K]);Aqk/Akk始终返回,固定为 head-major;gk/w/u/qg/kg/v_new是供反向使用的 head-major 中间量;- 公开
h固定为 sequence-major;内部hCompute保持 head-major 供 Finalize 使用——这是两个生命周期不同的张量; - 第 12 个返回值是 Python 层对
initial_state的原对象透传,不是 aclnn 输出。
输出保留策略对齐 fla-orgchunk_kda_fwd(对应提交0f0f0c97af39343855b43bbbaddcedfda5cb9d77):
| 条件 | 返回 |
|---|---|
output_final_state=true | 返回final_state,否则为None |
use_gate_in_kernel=false或disable_recompute=true | 返回gk |
| 始终 | 返回Aqk/Akk |
disable_recompute=true | 返回w/u/qg/kg/v_new |
disable_recompute=true或return_intermediate_states=true | 返回h |
aclnn L2 层的写出规则
fla_npu.ops.ascendc.chunk_kda_fwd是 12 返回值的低层语义封装,不涉及 CP。aclnn L2 不接收output_final_state/disable_recompute/return_intermediate_states三个布尔属性,每个可选输出是否写出仅由对应输出指针是否为空决定:
w/u/qg/kg/v_new/h的 L0 阶段固定写内部 compute 张量,L2 仅在对应指针非空时通过ViewCopy导出;指针为空时这些中间量只保留前向内部生命周期;gkOut非空时直接复用为gkCompute,避免在目标场景额外复制整张 FP32 gate;- 内部
hCompute是 FwdH 到 Finalize 的必需 head-major 阶段结果;hOut为空时仍会创建hCompute,只是不作为第 11 个 Python 返回值公开;hOut非空时 L2 先写 head-major 临时输出,再在导出边界转为 sequence-major; finalStateOut != nullptr同时表示本次需要计算并写出最终状态。
属性与支持范围
算子属性如下:
| 名称 | 默认值 | 支持范围 |
|---|---|---|
layout | BSND | BSND/BNSD/TND/NTD |
scale | 必传 | 通常为K**-0.5 |
chunk_size | 64 | 64/128 |
output_final_state | false | bool |
safe_gate | false | bool |
lower_bound | -5.0 | safe raw gate 时[-5,0) |
use_gate_in_kernel | false | bool |
disable_recompute | false | bool |
return_intermediate_states | false | bool |
state_v_first | false | bool |
支持范围:
- 平台:A2(
ascend910b)、A3(ascend910_93)、Ascend 950PR & 950DT 系列(ascend950)——三个平台在 chunk_kda_fwd_def.cpp 中均注册了OpAICoreConfig; K/V为[16,256]内 16 的倍数;交付重点覆盖 K=128、V=128/256;chunk_size为 64/128;- TND/NTD 均支持多 head;
- 变长调用最多 1024 条逻辑序列,rank-4 变长输入要求 B=1。
这些约束在 torch 适配层 chunk_kda_fwd_torch_adpt.h 中有显式校验:layout必须是BSND/BNSD/TND/NTD且大写;chunk_size只能是 64 或 128;rank-3 与 rank-4 输入的维度必须与 layout 匹配;0 < H <= H_v <= 128且H_v % H == 0;K/V必须是 16 的倍数且不大于 256;q/k/v必须同为 FP16 或 BF16。
调用途径与 API 详解
算子共有四条调用路径:
| 路径 | 入口 |
|---|---|
| 稳定 Python | fla_npu.ops.ascendc.chunk_kda_fwd |
| aclnn | aclnnChunkKdaFwdGetWorkspaceSize/aclnnChunkKdaFwd |
| legacy | 显式加载后的torch.ops.npu.npu_chunk_kda_fwd |
| 受限直调样例 | torch.ops.ascend_ops.chunk_kda_fwd_direct |
其中“受限直调样例”仅覆盖 dense BNSD、K=128、V=128/256,并保留“调用方传入已累计 gk”的低层测试接口;公开顶层语义以稳定 Python / aclnn 接口为准。
Python 主入口
from fla_npu.ops.ascendc import chunk_kda_fwd outputs = chunk_kda_fwd( q, k, v, g, beta, scale, chunk_size, layout="BSND", initial_state=None, output_final_state=False, cu_seqlens=None, chunk_indices=None, safe_gate=False, lower_bound=None, use_gate_in_kernel=False, A_log=None, dt_bias=None, disable_recompute=False, return_intermediate_states=False, state_v_first=False, )返回 12 元组(attn_out, final_state, gk, Aqk, Akk, w, u, qg, kg, v_new, h, initial_state);可选输出在 Python 层返回None,Aqk/Akk始终存在,其余保留策略见“输出契约”一节。
aclnn 接口
aclnnStatus aclnnChunkKdaFwdGetWorkspaceSize( const aclTensor *q, const aclTensor *k, const aclTensor *v, const aclTensor *g, const aclTensor *beta, const aclTensor *aLogOptional, const aclTensor *dtBiasOptional, const aclTensor *initialStateOptional, const aclIntArray *cuSeqlensOptional, const aclIntArray *chunkIndicesOptional, const char *layout, double scale, int64_t chunkSize, bool safeGate, double lowerBound, bool useGateInKernel, bool stateVFirst, const aclTensor *attnOut, const aclTensor *finalStateOut, const aclTensor *gkOut, const aclTensor *aqkOut, const aclTensor *akkOut, const aclTensor *wOut, const aclTensor *uOut, const aclTensor *qgOut, const aclTensor *kgOut, const aclTensor *vNewOut, const aclTensor *hOut, uint64_t *workspaceSize, aclOpExecutor *executor); aclnnStatus aclnnChunkKdaFwd( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);aclnn L2 只描述张量与算法契约,不接收、也不解释 autograd 重计算策略:attnOut/aqkOut/akkOut是必选输出;finalStateOut/gkOut/wOut/uOut/qgOut/kgOut/vNewOut/hOut均为相互独立的可选输出。output_final_state/disable_recompute/return_intermediate_states只存在于 Python 与 legacy torch 包装层,由上层按 FLA 保留策略决定向 L2 传入哪些输出指针。
输入输出布局约定
layout只解释 q/k/v/g/beta 输入;输出固定为:
attnOut:BSND 或 TND;finalStateOut:[N,H_v,K,V],stateVFirst=true时为[N,H_v,V,K];gkOut/AqkOut/AkkOut/wOut/uOut/qgOut/kgOut/vNewOut:BNSD/NTD;hOut:dense 为[B,N_c,H_v,K,V],varlen 为[N_c,H_v,K,V];stateVFirst=true时交换末两维。
完整示例
import torch from fla_npu.ops.ascendc import chunk_kda_fwd B, T, H, K, V = 1, 128, 4, 128, 128 q = torch.randn(B, T, H, K, device="npu", dtype=torch.float16) k = torch.randn_like(q) v = torch.randn(B, T, H, V, device="npu", dtype=torch.float16) g = -torch.rand(B, T, H, K, device="npu", dtype=torch.float32) * 0.01 beta = torch.rand(B, T, H, device="npu", dtype=torch.float32) attn_out, final_state, *_ = chunk_kda_fwd( q, k, v, g, beta, K ** -0.5, 64, layout="BSND", output_final_state=True, safe_gate=True, ) assert attn_out.shape == (B, T, H, V) assert final_state.shape == (B, H, K, V)L2 调度与内核四阶段设计
L2 调度流
raw g -> ChunkKdaFwd[ gate cumsum -> Prepare/Post-WU -> FwdH -> Finalize ] -> attn_outaclnnChunkKdaFwd在 L2 层完成公开 layout 的连续化和必要视图转换。A5 的 BF16、chunk=64、K=V=128 dense 对齐快路径保持单次物理ChunkKdaFwdL0;A5 其他多 chunk 场景将同一个私有 L0 按 Gate/Prepare、Post-WU、FwdH、Finalize 四个阶段依次提交,使阶段间通过物理 launch 边界重置事件状态。A2/A3 与单 chunk 场景仍使用单次物理 L0。阶段选择仅使用私有stage属性,不增加公开属性、接口字段或独立算子原型——这一设计在 op_kernel/chunk_kda_fwd.cpp 中有直接对应:KDA_STAGE_FULL = -1、KDA_STAGE_GATE_PREPARE = 0、KDA_STAGE_POST_WU = 1、KDA_STAGE_FWD_H = 2、KDA_STAGE_FINALIZE = 3,DispatchStage依据tiling.stage分发到对应阶段实现。
KdaGateCumsum
将 raw/已激活 gate 转为 FP32 chunk-local log2 累计值:
gk = cumsum(gate) / ln(2)该阶段同时保留独立 L2 接口供 GDN2 调用(对应仓库中独立的 kda_gate_cumsum 目录),输入输出固定为 BNSD/NTD。
Prepare
只读取q/k/v/gk/beta及变长元数据,产生:
Aqk, Akk, qg, qg_scaled, w_seed, u_seed矩阵计算和三角求逆使用FP32 累积;公开中间量在写回时转为 q dtype。从内核侧看,RunChunkKdaPrepare在 chunk_kda_fwd.cpp 中以SAFE_GATE、T、float、BETA_T为模板参数实例化,即内部统一以 FP32 参与数值主计算。
Post-WU
只读取k/gk/w_seed/Akk/u_seed,产生:
w, u, kg, v_new_seedAkk的 head 循环按H_v执行,GQA 映射只在读取 q/k head 时换算,避免按H_k重复或漏算——这是 GQA 场景下保证正确性与性能的关键实现细节。
FwdH state propagation
读取kg/w/u/gk和可选initial_state,计算 chunk 间递推:
v_new = u - w @ h_prev h_next = exp2(gk_last) * h_prev + kg^T @ v_newarch35 路径复用与ChunkGatedDeltaRuleFwdH(见仓库 chunk_gated_delta_rule_fwd_h)相同的数学实现;其他场景在ChunkKdaFwd内嵌共享 FwdH 实现。独立 GDN L0 原型继续保留给其他调用方,key-wisegk固定使用exp2。内核入口按tiling.vHeadDim > 128选择GDNFwdHTileShapes256或GDNFwdHTileShapes128两个 tile shape(见 chunk_kda_fwd.cpp)。
Finalize
只读取qg_scaled/Aqk/v_new/h,计算:
attn_out = qg_scaled @ h + Aqk @ v_newkernel 内直接按 BSND/TND 写出attn_out;供反向使用的中间量保持 BNSD/NTD。
状态布局与重计算策略
内部递推统一使用[...,K,V]。state_v_first=true时,L2 在进入 FwdH 前转置 initial state。内部hCompute始终保持 head-major 供 Finalize 消费;公开hOut在 L2 导出边界转为 sequence-major,并按state_v_first决定末两维顺序。final_state按序列排列,与 FLA 顶层输出一致。
重计算策略上,L2 不理解 autograd 重计算策略,final_state/gk/w/u/qg/kg/v_new/h是相互独立的OPTIONAL_OUTPUT:非空指针表示导出,空指针表示不公开该结果。单 launch 路径为隐藏输出传递固定 ABI 占位,并由 tiling 在 kernel workspace 中承接实际中间结果;A5 四段 launch 路径将阶段间依赖的gk/w/u/qg/kg/v_new/h/final_state和私有qg_scaled/u_seed物化为 executor 内部张量,使后续 launch 不依赖前一 launch 的 kernel workspace。公开输出存在时直接作为内部目标使用。Python/legacy 包装层对齐 fla-orgchunk_kda_fwd提交0f0f0c97af39343855b43bbbaddcedfda5cb9d77的保留规则:disable_recompute=false时不保留w/u/qg/kg/v_new;disable_recompute=true或return_intermediate_states=true时保留公开hOut;use_gate_in_kernel=false或disable_recompute=true时保留gk;final_state只在output_final_state=true时创建公开输出。
模板化方案与 tiling key
ChunkKdaFwd只有一个外层 op_kernel/chunk_kda_fwd.cpp 入口(extern "C" __global__ __aicore__ void chunk_kda_fwd)和一个私有 L0 类型。A5 实现位于 op_kernel/arch35/*.h,host 侧 A5 模板选择位于 op_host/arch35/chunk_kda_fwd_tiling_impl.h。Prepare、Post-WU、Finalize 的内部实现头与统一 kernel 入口同属chunk_kda_fwd/op_kernel/目录,不存在对应的独立 L0 原型或.cpp入口。A5 四段路径只是用不同私有stage属性连续调用该入口。
两个编译期 tiling key 是同一 L0 的场景变体,不是平台编号、独立算子或独立接口:
tiling key=1:非 chunk=64、K=V=128 场景的通用模板族;tiling key=2:chunk=64、K=V=128 模板族,包括 dense、tail 和 varlen。
A2/A3/A5 均生成两个 key;同一个 key 内再由编译架构选择根目录通用实现或arch35/实现。host 的SetTilingKey只检查 chunk、K、V,不检查 SoC。在 arch35 上,key2 的 dense 对齐场景使用单 launch 和 arch35 FwdH;融合 score 写回在跳过共享 PostWU 时会额外物化以块尾 gate 为参考的最终kg,供 FwdH 和可选公开输出共同使用。A5 多 chunk 的 tail/varlen 以及 key1 泛化场景使用四段 launch。tiling key 与私有stage均不改变公开算子原型、输出契约或数学定义。
tiling 侧(chunk_kda_fwd_tiling.cpp)为 workspace 规划了固定 ABI:512 字节对齐(KDA_ALIGN = 512),包含三角求逆 scratch(KDA_SOLVE_SCRATCH_SLOTS = 5,流水深度 4)、score 队列(KDA_SCORE_QUEUE_SLOTS = 4,KDA_SCORE_SCRATCH_PLANES = 3)与 GDN 流水(KDA_GDN_PIPELINE_DEPTH = 2)等区域,并通过HasOutput按实例输出 shape 判断各可选输出是否真正需要导出。
性能设计要点
设计文档明确了以下性能手段:
- Prepare 的右矩阵在L1 驻留,避免 K/K^T 重复搬运和重复转置;
- AIC 使用 L1/L0 双缓冲组织 MTE2、MTE1、Cube、Fixpipe 流水;
- AIV 使用输入 staging ping-pong,使下一 tile 的 MTE2 与当前 tile 的 VEC 重叠;
- A5 VEC 路径使用 regbase 双发射特化,数值主计算仍保持 FP32;
- inter-sub-chunk 合并使用独立 workspace 区域,避免阻塞主 tile 流水。
性能结论只使用msopprof评测;目标回归 case 定义在tests/op_cases/chunk_kda_fwd.json(该路径位于算子开发工程的测试目录,数值测试对应tests/operators/chunk_kda_fwd/accuracy/,性能测试对应tests/operators/chunk_kda_fwd/performance/profile.py+msopprof)。
验证矩阵
算子验证覆盖如下组合:
- 平台:A2/A3/A5;
- dtype:FP16/BF16;
- layout:BSND/BNSD/TND/NTD;
- gate:raw/已激活、safe true/false;
- Shape:K=128,V=128/256,chunk=64/128,dense/varlen/tail/GQA;
- 属性:final state、重计算策略、
state_v_first。
唯一用例规格是tests/op_cases/chunk_kda_fwd.json;数值测试位于tests/operators/chunk_kda_fwd/accuracy/,性能使用tests/operators/chunk_kda_fwd/performance/profile.py与msopprof评测。
相关文件索引
- 算子主文档:csrc/attention/chunk_kda_fwd/README.md
- API 文档:csrc/attention/chunk_kda_fwd/docs/api.md
- 设计文档:csrc/attention/chunk_kda_fwd/docs/design.md
- 内核统一入口与阶段分发:op_kernel/chunk_kda_fwd.cpp
- 内核公共头/变长支持:op_kernel/chunk_kda_fwd_common.h、op_kernel/chunk_kda_fwd_varlen.h
- arch35 快路径实现:op_kernel/arch35/
- 算子原型定义(三平台注册):op_host/chunk_kda_fwd_def.cpp
- tiling 与 workspace 规划:op_host/chunk_kda_fwd_tiling.cpp
- torch 适配层与参数校验:chunk_kda_fwd_torch_adpt.h
- vLLM 侧真实调用示例:vllm_ascend/ops/kda.py
- 关联的独立 gate cumsum 算子:csrc/attention/kda_gate_cumsum
- 人工智能
- 大模型
- 模型推理服务
- Ascend
- CANN
【免费下载链接】vllm-ascend
Community maintained hardware plugin for vLLM on Huawei Ascend
相关推荐
vllm-ascend ChunkKdaFwd 算子 API 深度解析:Python 入口、aclnn 契约与 Gate 数值语义
vllm ascend ChunkKdaFwd 算子 API 深度解析:Python 入口、aclnn 契约与 Gate 数值语义 本篇指南以 vllm asc
人工智能大模型模型推理服务AscendCANNvllm-ascend ChunkKdaFwd 算子设计解析:gate 线性注意力前向的 L2 调度、阶段流水与昇腾模板化实现
vllm ascend ChunkKdaFwd 算子设计解析:gate 线性注意力前向的 L2 调度、阶段流水与昇腾模板化实现 本篇技术指南围绕 vllm as
人工智能大模型模型推理服务AscendCANNvllm-ascend 自定义算子 MlaPrologV3K3 全通路 API 指南:从 torch 单算子入口到 aclnn 与 Ascend C 直调
vllm ascend 自定义算子 MlaPrologV3K3 全通路 API 指南:从 torch 单算子入口到 aclnn 与 Ascend C 直调 本篇
人工智能大模型模型推理服务AscendCANN
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考