资讯详情

CANN ops-transformer QkvRmsNormRopeCache 算子详解:融合 QKV 拆分、RmsNorm、RoPE 与量化 KV Cache 写入

📅 2026/9/18 8:45:49 | 华诺云谱 👁 阅读
CANN ops-transformer QkvRmsNormRopeCache 算子详解:融合 QKV 拆分、RmsNorm、RoPE 与量化 KV Cache 写入
CANN ops-transformer QkvRmsNormRopeCache 算子详解融合 QKV 拆分、RmsNorm、RoPE 与量化 KV Cache 写入【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本文基于 CANN ops-transformer 开源仓库中 posembedding/qkv_rms_norm_rope_cache/README.md 及其配套的接口文档、Host 侧实现、Kernel 实现与测试代码整理而成。QkvRmsNormRopeCache 将 Transformer 解码阶段 QKV 投影后到 KV Cache 写入之间的一系列算子SplitVD 拆分、RmsNorm、RoPE 旋转位置编码、可选量化、按 index 的 Scatter 写入融合为单个 NPU 算子是 PagedAttention 类长序列推理场景中典型的访存密集型融合算子。读完本文你将掌握该算子的计算语义、输入输出约束、属性配置、aclnn 两段式调用方法以及其 Tiling 与 Kernel 的底层实现思路。功能定位一次算子调用完成 QKV 处理 KV Cache 写入QkvRmsNormRopeCache 的输入是 QKV 融合张量即已经过权重投影拼接的 qkv 矩阵算子内部依次完成SplitVD将融合张量按头维度切分为 q、k、v 三个分量RmsNorm对 q、k 分量沿最后一维做均方根归一化v 分量不参与RoPEHalf-and-Half 旋转位置编码对完成 RmsNorm 的 q、k 施加旋转位置编码Quant可选按 k_scale/v_scale对称量化或叠加 k_offset/v_offset非对称量化将 k、v 量化为 INT8Scatter根据 index 将结果写入预先申请的 PA_NZ 格式 KV CachePagedAttention 的分页缓存中。最终输出为q_out、k_cache、v_cache以及可选的q_out_before_quant、k_out_before_quant、v_out_before_quant量化与 Scatter 之前的中间结果。产品支持情况产品是否支持Ascend 950PR/Ascend 950DT×Atlas A3 训练系列产品/Atlas A3 推理系列产品√Atlas A2 训练系列产品/Atlas A2 推理系列产品√Atlas 200I/500 A2 推理产品×Atlas 推理系列产品×Atlas 训练系列产品×与算子定义文件中的 AICore 配置一致在 qkv_rms_norm_rope_cache_def.cpp 中通过this-AICore().AddConfig(ascend910b)与AddConfig(ascend910_93)注册了对应平台的 AICore 版本其中 ascend910b 对应 Atlas A2 系列ascend910_93 对应 Atlas A3 系列。支持的场景类型当前版本仅支持单一融合场景场景类型情况概要cache_mode 为 PA_NZq 无量化k 和 v 支持无量化、对称量化和非对称量化q_out_before_quant/k_out_before_quant/v_out_before_quant 不输出qkv Shape 为 [B_qkv * S_qkv, N_qkv * D_qkv]q、k、v 具有完全相同的 D 维度。计算与输出对应关系qkv → SplitVD → q、k、vq → RmsNorm、RoPE → q_outk → RmsNorm、RoPE、Quant(可选)、Scatter → k_cachev → Quant(可选)、Scatter → v_cache需要说明的是虽然算子定义OpDef与 aclnn 接口均声明了可选的 before_quant 输出由is_output_qkv属性控制但当前成文的支持场景以 README 声明的表格为准即不输出 before_quant 的 PA_NZ 场景。从 Kernel 入口 qkv_rms_norm_rope_cache.cpp 可以看到当前仅编译了 TilingKey3对应 B16 PA_NZ Quant这一个 Kernel 分支。计算原理与数学公式以下公式均来自 README 与 aclnn 接口文档并已在 tests/st/aclnnQkvRmsNormRopeCache/executor_aclnnQkvRmsNormRopeCache.py 的qkvRmsNormRopeCache_golden参考实现中逐一落地可作为精度对齐的依据。(1) SplitVDQKV 融合张量切分设 N_q、N_k、N_v 分别为 q、k、v 分量的注意力头数量必须满足$$ \begin{cases} N_k N_v \ N_{qkv} N_k N_v N_q \ D_{qkv} D_q D_k D_v \end{cases} $$切分公式按 head 维度切片$$ \begin{aligned} q qkv[..., [:N_q] \times D_{qkv}] \ k qkv[..., [N_q:-N_v] \times D_{qkv}] \ v qkv[..., [-N_v:] \times D_{qkv}] \end{aligned} $$在 Host 侧 Tiling 校验中该约束被严格检查numHead_ ! numHeadQ_ numHeadK_ numHeadV_或numHeadK_ ! numHeadV_时直接返回失败且额外要求numHeadQ_ % numHeadK_ 0见 qkv_rms_norm_rope_cache_tiling.cpp。(2) RmsNorm均方根归一化RmsNorm 沿最后一维feature dimension进行通用于 q、k 分量。设 x 为输入、y 为输出、gamma 为可学习缩放参数$$ squareX x \times x $$$$ meanSquareX squareX.mean(dim -1, keepdim True) $$$$ rms \sqrt{meanSquareX epsilon} $$$$ y (x / rms) \times gamma $$epsilon用于防止除 0默认值 1e-6。在 Tiling 实现中Host 侧预先计算reciprocal_ 1.0 / qkvDim_供 Kernel 做均值运算qkv_rms_norm_rope_cache_tiling.cpp。(3) RoPEHalf-and-Half 旋转位置编码此处的 y 指代完成 RmsNorm 的输出d 为特征维度$$ y1 y[..., :d/2] $$$$ y2 y[..., d/2:] $$$$ y_RoPE torch.cat((-y2, y1), dim -1) $$$$ y_embed (y \times cos) y_RoPE \times sin $$其中 cos、sin 为预先计算好的位置编码表shape 为 [B_qkv * S_qkv, D_rope]D_rope D_qkv。(4) Quantk/v 量化可选无量化k_cache/v_cache 保持 FLOAT16/BFLOAT16$$ kQuant kRoPE $$$$ vQuant v $$对称量化k_cache/v_cache 为 INT8需要 k_scale/v_scale$$ kQuant kRoPE / kScale $$$$ vQuant v / vScale $$非对称量化INT8需要 k_scale/v_scale 与 k_offset/v_offset$$ kQuant kRoPE / kScale kOffset $$$$ vQuant v / vScale vOffset $$注意q 分量当前不支持量化量化与否由 k_cache/v_cache 的张量数据类型决定——为 INT8 即量化为与 qkv 相同的 FLOAT16/BFLOAT16 即不量化。Host 侧校验强制 k_cache/v_cache 数据类型必须二选一且 k_scale 与 k 量化状态必须严格对应qkv_rms_norm_rope_cache_tiling.cpp。参数说明算子aclnn 接口的完整参数如下表其中输入/输出列同时标注了参数在接口中的角色参数名输入/输出/属性描述数据类型数据格式qkv输入用于切分出 q、k、v 的输入数据对应公式中的 qkv。shape 为 [B_qkv * S_qkv, N_qkv * D_qkv]FLOAT16、BFLOAT16NDq_gamma输入用于 q 的 rms_norm 计算的输入数据对应公式中的 gamma。与输入 qkv 数据类型相同shape 为 [D_qkv]FLOAT16、BFLOAT16NDk_gamma输入用于 k 的 rms_norm 计算的输入数据对应公式中的 gamma。与输入 qkv 数据类型相同shape 为 [D_qkv]FLOAT16、BFLOAT16NDcos输入用于 rope 计算的输入数据对输入张量进行余弦变换对应公式中的 cos。与输入 qkv 数据类型相同shape 为 [B_qkv * S_qkv, 1 * D_rope]D_rope D_qkv 且要求D_qkv * qkv 数据类型所占字节数可被 32 整除FLOAT16、BFLOAT16NDsin输入用于 rope 计算的输入数据对输入张量进行正弦变换对应公式中的 sin。与输入 cos 的数据类型、格式保持一致FLOAT16、BFLOAT16NDindex输入用于指定写入 cache 的具体索引位置。shape 为 [B_qkv * S_qkv]INT64NDq_out输入/输出提前申请的 cache输入输出同地址复用in-place。与输入 qkv 数据类型相同shape 为 [B_qkv * S_qkv, N_q * D_qkv]FLOAT16、BFLOAT16NDk_cache输入/输出提前申请的 cache输入输出同地址复用。与输入 qkv 数据类型相同k 不量化或 INT8k 量化。shape 为 [BlockNum, N_k * D_qkv // 16, BlockSize, 16]不量化或 [BlockNum, N_k * D_qkv // 32, BlockSize, 32]量化FLOAT16、BFLOAT16、INT8NDv_cache输入/输出提前申请的 cache输入输出同地址复用。与输入 qkv 数据类型相同v 不量化或 INT8v 量化。shape 为 [BlockNum, N_k * D_qkv // 16, BlockSize, 16]不量化或 [BlockNum, N_k * D_qkv // 32, BlockSize, 32]量化FLOAT16、BFLOAT16、INT8NDk_scale可选输入当 k_cache 数据类型为 INT8 时需要此输入参数对应公式中的 kScale。shape 为 [N_k, D_qkv]FLOAT32NDv_scale可选输入当 v_cache 数据类型为 INT8 时需要此输入参数对应公式中的 vScale。shape 为 [N_v, D_qkv]FLOAT32NDk_offset可选输入当 k_cache 数据类型为 INT8 且 k_scale 输入存在并量化场景为非对称量化时需要此参数输入对应公式中的 kOffset。shape 为 [N_k, D_qkv]FLOAT32NDv_offset可选输入当 v_cache 数据类型为 INT8 且 v_scale 输入存在并量化场景为非对称量化时需要此参数输入对应公式中的 vOffset。shape 为 [N_v, D_qkv]FLOAT32NDq_out_before_quant可选输出即将写入到 q_out 中的数据FLOAT16、BFLOAT16NDk_out_before_quant可选输出即将写入到 k_cache 中的数据未经量化和 Scatter 前的中间计算结果FLOAT16、BFLOAT16NDv_out_before_quant可选输出即将写入到 v_cache 中的数据未经量化和 Scatter 前的中间计算结果FLOAT16、BFLOAT16NDqkv_size属性按 [B_qkv, S_qkv, N_qkv, D_qkv] 顺序传入提供 qkv 矩阵的 B、S、N、D 维度尺寸INT64-head_nums属性按 [N_q, N_k, N_v] 顺序传入提供 qkv 矩阵中 qkv 分量单元的 N 维度尺寸INT64-epsilon可选属性用于防止 RmsNorm 计算除 0 错误对应公式中的 epsilon默认值为 1e-6FLOAT32-cache_mode可选属性cache 格式的选择标记目前只支持 PA_NZ默认值为 PA_NZCHAR*-is_output_qkv可选属性表示是否需要输出各 cache 输出中对应内容在未经量化和 Scatter 前的原始值默认值为 falseBOOL-上述参数约束均可在源码中找到对应校验qkv 数据类型仅支持 FLOAT16/BFLOAT16qkv_rms_norm_rope_cache_tiling.cppgamma 维度必须为 [D_qkv] 且与 qkv 同 dtypecos/sin 的 S 维度允许为 B_qkv * S_qkv 或 B_qkv 两种广播形态rope_seq_由此推导见 CheckCosSinValidindex 必须为 INT64 且元素个数等于 B_qkv * S_qkv。属性默认值在代码中的体现算子定义 qkv_rms_norm_rope_cache_def.cpp 中qkv_size、head_nums为必选 ListInt 属性epsilon可选默认1e-6fcache_mode可选默认字符串PA_NZis_output_qkv可选默认false。Tiling 侧读取时同样做了空指针兜底qkv_rms_norm_rope_cache_tiling.cpp并将 cache_mode 通过{PA_NZ: CacheMode::PA_NZ}映射表校验不支持的模式直接报错。约束说明输入 shape 限制B_qkv 为输入 qkv 的 batch_sizeS_qkv 为 sequence length大小由 qkv_size 决定N_qkv 为输入 qkv 的 head number。D_qkv 为 head dim目前仅支持 128Tiling 中OP_CHECK_IF(qkvDim_ ! 128, ...)直接校验见 qkv_rms_norm_rope_cache_tiling.cpp。D_q D_k D_v D_qkv且D_qkv * qkv 数据类型所占字节数可被 32 整除根据 rope 规则D_k 和 D_q 必须为偶数。cache_mode 为 PA_NZ 场景下D_k、D_q 需 32B 对齐BlockSize 需 32B 对齐32B 对齐的具体值由 cache 数据类型决定以 BlockSize 为例cache 为 int8 时需满足 BlockSize % 32 0cache 为 float16 时需满足 BlockSize % 16 0若 k_cache 与 v_cache 的 dtype 不一致BlockSize 需同时满足 BlockSize % 32 0 和 BlockSize % 16 0。对应常量可在 qkv_rms_norm_rope_cache_tiling.h 中找到INT8_BLOCK_ALIGN_NUM 32、FP16_BLOCK_ALIGN_NUM 16并在 CheckKCacheValid 中校验 k_cache 的中间两维与对齐关系BlockNum 为写入 cache 的内存块数大小由用户输入场景决定要求BlockNum Ceil(S_qkv / BlockSize) * B_qkvTiling 中有对应校验GM 空间限制设 requireMemory 为存放数据所需的空间大小需满足 requireMemory (B_qkv * S_qkv * N_qkv * D_qkv 2 * D_qkv 2 * B_qkv * S_qkv * D_qkv B_qkv * S_qkv * N_q * D_qkv BlockNum * BlockSize * N_v * D_qkv BlockNum * BlockSize * N_k * D_qkv) * sizeof(FLOAT16) B_qkv * S_qkv * sizeof(INT64) (2 * N_k * D_qkv 2 * N_v) * sizeof(FLOAT) 当计算出的 requireMemory 超过当前 AI 处理器的 GM 空间总大小时不支持使用该算子。其他限制对 indexvalue 值范围为 [-1, BlockNum * BlockSize)value 数值不可以重复index 为 -1 时代表跳过更新该跳过写入语义与 PagedAttention 中 padding token 的占位行为对应k_scale、v_scale 表示对称量化的缩放因子若传参则值不能为 0aclnn 接口默认确定性实现deterministic。aclnn 两段式接口调用QkvRmsNormRopeCache 算子通过 aclnn 接口调用遵循 CANN 算子库通用的两段式接口模式详见接口文档 aclnnQkvRmsNormRopeCache.md必须先调用第一段aclnnQkvRmsNormRopeCacheGetWorkspaceSize完成参数校验、构图并计算 workspace 大小再调用第二段aclnnQkvRmsNormRopeCache真正执行计算。函数原型aclnnStatus aclnnQkvRmsNormRopeCacheGetWorkspaceSize( const aclTensor *qkv, const aclTensor *qGamma, const aclTensor *kGamma, const aclTensor *cos, const aclTensor *sin, const aclTensor *index, aclTensor *qOut, aclTensor *kCache, aclTensor *vCache, const aclTensor *kScaleOptional, const aclTensor *vScaleOptional, const aclTensor *kOffsetOptional, const aclTensor *vOffsetOptional, const aclIntArray *qkvSize, const aclIntArray *headNums, double epsilon, char *cacheModeOptional, const aclTensor *qOutBeforeQuant, const aclTensor *kOutBeforeQuant, const aclTensor *vOutBeforeQuant, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnQkvRmsNormRopeCache( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)接口实现位于 aclnn_qkv_rms_norm_rope_cache.cpp其中几个关键实现细节值得注意空指针检查CheckNotNull对 qkv、qGamma、kGamma、cos、sin、index、qOut、kCache、vCache、qkvSize、headNums 共 11 个必传参数做空指针校验返回ACLNN_ERR_PARAM_NULLPTR错误码 161001Contiguous 处理所有输入含可选输入先经l0op::Contiguous转换为连续张量非连续 Tensor 也可正常调用可选输出与is_output_qkv的联动bool isOutputQkv qOutBeforeQuant nullptr ? false : true即是否输出 before_quant 中间结果完全由 qOutBeforeQuant 是否为空决定三个 before_quant 输出必须同时提供或同时为空否则返回ACLNN_ERR_INNER_NULLPTRaclnn_qkv_rms_norm_rope_cache.cppworkspace 获取*workspaceSize uniqueExecutor-GetWorkspaceSize()从构建好的执行器中取得随后通过ReleaseTo(executor)把执行器交给第二段接口。返回值与错误码第一段接口完成入参校验出现以下场景时报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001传入的 qkv、qGamma、kGamma、cos、sin 等为空指针ACLNN_ERR_PARAM_INVALID161002输入或输出的数据类型不在支持范围内输入或输出的参数维度不在支持范围内dim 不在指定的取值范围内完整调用示例仓库提供了可直接参考的完整 C 样例 test_aclnn_qkv_rms_norm_rope_cache.cppUT 目录下还有一份用于单测的副本 tests/ut/op_host/op_api/test_aclnn_qkv_rms_norm_rope_cache.cpp。下面提炼其关键流程与张量形状设计#include acl/acl.h #include aclnnop/aclnn_qkv_rms_norm_rope_cache.h // 样例场景B16, S3, Nqkv18(Nq16, Nk1, Nv1), D128 // 即 qkv [48, 2304]q_out [48, 2048] // kCache/vCache 为 PA_NZ 分页缓存 [BlockNum16, Nk*D/324, BlockSize128, 32]量化 INT8 场景 std::vectorint64_t qkvShape {48, 2304}; std::vectorint64_t qGammaShape {128}; std::vectorint64_t kGammaShape {128}; std::vectorint64_t cosShape {48, 128}; std::vectorint64_t sinShape {48, 128}; std::vectorint64_t indexShape {48}; std::vectorint64_t qOutShape {48, 2048}; std::vectorint64_t kCacheShape {16, 4, 128, 32}; std::vectorint64_t vCacheShape {16, 4, 128, 32}; std::vectorint64_t kScaleShape {1, 128}; std::vectorint64_t vScaleShape {1, 128}; std::vectorint64_t qkv_size_list {16, 3, 18, 128}; // [B, S, N, D] std::vectorint64_t head_nums_list {16, 1, 1}; // [Nq, Nk, Nv]调用主流程分七步int main() { // 1. device/stream 初始化 int32_t deviceId 0; aclrtStream stream; auto ret Init(deviceId, stream); // 2. 通过 aclrtMalloc aclrtMemcpy 构造 device 侧张量 // 用 aclCreateTensor 创建 aclTensor详见 CreateAclTensor 辅助函数 // 本例 kCache/vCache 使用 ACL_INT8量化kScale/vScale 使用 ACL_FLOAT // 其余输入输出均使用 ACL_FLOAT16index 使用 ACL_INT64 // 3. 第一段接口计算 workspace 并构建 executor uint64_t workspaceSize 0; aclOpExecutor *executor; ret aclnnQkvRmsNormRopeCacheGetWorkspaceSize( qkv, qGamma, kGamma, cos, sin, index, qOut, kCache, vCache, kScale, vScale, // k/v 量化 scale nullptr, nullptr, // k/v offset对称量化场景传空 qkv_size, head_nums, epsilon /* 1e-6 */, cacheMode /* PA_NZ */, nullptr, nullptr, nullptr, // 三个 before_quant 输出不输出则传空 workspaceSize, executor); // 按需申请 workspace void *workspaceAddr nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); } // 4. 第二段接口执行计算 ret aclnnQkvRmsNormRopeCache(workspaceAddr, workspaceSize, executor, stream); // 5. 同步等待任务执行结束 ret aclrtSynchronizeStream(stream); // 6. 将结果从 device 拷贝回 hostPrintOutResult并 aclDestroyTensor 释放 aclTensor // 7. aclrtFree 释放 device 内存与 workspaceaclrtDestroyStream / aclrtResetDevice / aclFinalize 收尾 }该样例是一个典型的 INT8 对称量化 PA_NZ 分页缓存场景q 不做量化直接输出到 q_outk、v 经 RmsNorm/RoPE 与量化后按 index 散射写入 [16, 4, 128, 32] 的 kCache/vCache。index 的取值范围与跳过更新语义可在 ST 测试的输入构造逻辑中找到对应实现executor_aclnnQkvRmsNormRopeCache.py从[-1, BlockNum * BlockSize)中无放回抽样超出部分用 -1 填充验证了index 不可重复、-1 表示跳过写入的约束。源码级实现解析Tiling多路并行切分与 UB 空间预算Host 侧 Tiling 的核心在QkvRmsNormRopeCacheTilingDs::CalUbTiling()qkv_rms_norm_rope_cache_tiling.cpp其切分策略可以概括为按 Q/K/V 三路分配 AI CorecoreNumKv按 (NkNv)/Nqkv 的比例估算 KV 路核数再按FAC_K 0.6的固定系数在 K 路与 V 路之间分配K 路因含 RmsNorm 与 RoPE 计算量更大见 qkv_rms_norm_rope_cache_tiling.h 中FAC_K的注释剩余核数分给 Q 路按 token 总数均分每路各自计算blockFactor每个核处理的 token 数与blockDim该路实际使用的核数最终blockDim blockDimQ blockDimK blockDimVUB 切分因子K 路需要 inQueueX/inQueueY 双缓冲、4 个 locBuf中间结果缓冲与 outQueueV 路计算量较小buffer 预算更紧凑。ubFactor通过(ubSize - UB_RESERVED_BYTES) / spaceWithUbfactor计算UB_RESERVED_BYTES 预留 1KB 冗余写回 tilingDatabatchSize、seqLength、head 数、blockNum/blockSize、epsilon、各路 blockFactor/blockDim/ubFactor、reciprocal、isOutputQkv、isKQuant、isVQuant等全部字段写入 tiling 结构体qkv_rms_norm_rope_cache_tiling.h供 Kernel 侧读取TilingKey 生成tilingKey cacheMode 1PA_NZ 对应 key3与 Kernel 入口编译的QKV_RMS_NORM_ROPE_CACHE_B16_PA_NZ_QUANT 3分支对应。KernelAIV 上的融合计算Kernel 入口 qkv_rms_norm_rope_cache.cpp 是一个 AIV-only 的 AI Core 程序extern C __global__ __aicore__ void qkv_rms_norm_rope_cache( GM_ADDR qkv, GM_ADDR q_gamma, GM_ADDR k_gamma, GM_ADDR cos, GM_ADDR sin, GM_ADDR index, GM_ADDR q_out, GM_ADDR k_cache, GM_ADDR v_cache, GM_ADDR k_scale, GM_ADDR v_scale, GM_ADDR k_offset, GM_ADDR v_offset, GM_ADDR q_out_out, GM_ADDR k_cache_out, GM_ADDR v_cache_out, GM_ADDR q_out_proto, GM_ADDR k_cache_proto, GM_ADDR v_cache_proto, GM_ADDR workspace, GM_ADDR tiling) { KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_AIV_ONLY); TPipe pipe; if (TILING_KEY_IS(QKV_RMS_NORM_ROPE_CACHE_B16_PA_NZ_QUANT)) { GET_TILING_DATA_WITH_STRUCT(QkvRmsNormRopeCacheTilingData, tiling_data_in, tiling); const QkvRmsNormRopeCacheTilingData* __restrict tilingData tiling_data_in; KernelQkvRmsNormRopeCacheQuantB16PANZDTYPE_QKV, DTYPE_K_CACHE, DTYPE_V_CACHE op(pipe, tilingData); op.Init(qkv, q_gamma, k_gamma, cos, sin, index, q_out, k_cache, v_cache, k_scale, v_scale, k_offset, v_offset, q_out_proto, k_cache_proto, v_cache_proto); op.Process(); } }Kernel 主体KernelQkvRmsNormRopeCacheQuantB16PANZ位于 qkv_rms_norm_rope_cache_b16_pa_nz_quant.hFP16/BF16 数据通路PA_NZ 分页缓存布局支持 k/v 量化公共工具与常量定义在 qkv_rms_norm_rope_cache_comm.h。模板参数 DTYPE_QKV、DTYPE_K_CACHE、DTYPE_V_CACHE 在编译期实例化不同的量化组合不量化 / K 量化 / V 量化与 qkv_rms_norm_rope_cache_binary.json 中注册的 fp16、bf16、fp16KquantVquant、bf16KquantVquant 等 bin 组合一一对应。InferShapein-place 输出与可选输出InferShape 实现见 qkv_rms_norm_rope_cache_infershape.cpp要点q_out/k_cache/v_cache 三个主输出直接复用对应输入的 shapein-place 地址复用is_output_qkv为 true 时推导三个 before_quant 可选输出的 shapeqOutProto 复用 qOut 输入 shapek/vCacheProto 复用 qkv 输入 shape并把 kCache 的中间两维还原为 [N_k * D_qkv]dim1 * dim3即把 PA_NZ 块布局还原为逻辑形状未知 rank-2 动态 shape场景将可选输出置为全 -1。测试与验证仓库围绕该算子提供了三级测试ST 系统测试tests/st/aclnnQkvRmsNormRopeCache/下包含 ATK 测试框架的用例配置文件 atk_aclnnQkvRmsNormRopeCache.json 与执行脚本 executor_aclnnQkvRmsNormRopeCache.py。脚本中的qkvRmsNormRopeCache_golden以 PyTorch 实现 SplitVD、RmsNorm、RoPE、量化与按 index 写入的参考逻辑用于与 NPU 实际输出做精度比对Host 侧 UTtests/ut/op_host/下覆盖 infershape 校验test_QkvRmsNormRopeCache_infershape.cpp、tiling 计算test_qkv_rms_norm_rope_cache_tiling.cpp与 aclnn 接口调用test_aclnn_qkv_rms_norm_rope_cache.cppKernel 侧 UTtests/ut/op_kernel/下提供 Kernel 级别的单元测试test_qkv_rms_norm_rope_cache.cpp。小结与适用前提QkvRmsNormRopeCache 是 CANN ops-transformer 中面向 PagedAttention 推理优化的典型融合算子把QKV 拆分 → RmsNorm → RoPE → 量化 → 分页写入五步合并为一次 NPU 计算显著减少中间张量的 GM 往返。使用前请重点确认平台Atlas A2/A3 训练与推理系列ascend910b / ascend910_93维度D_qkv 固定为 128D_q/D_k 为偶数且 32B 对齐BlockSize 按 cache dtype 满足 16/32 对齐数据流q 无量化k/v 的量化与否由 k_cache/v_cache 的 dtypeINT8 或与 qkv 相同决定量化时必须配套传 scale非对称再加 offsetindex 语义取值 [-1, BlockNum * BlockSize)不可重复-1 表示跳过该位置的写入调用方式严格遵循 aclnn 两段式接口先用 GetWorkspaceSize 取 workspace 大小并构建 executor再执行计算。如需继续深入可对照阅读仓库中的 接口文档、算子定义、Tiling 实现 与 Kernel 实现 完成端到端理解。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。