CANN 算子性能调优实战:KvRmsNormRopeCache 从 MemBase 到 RegBase 的平滑迁移与寄存器级 VF 计算
2026/9/18 20:58:41 网站建设 项目流程

CANN 算子性能调优实战:KvRmsNormRopeCache 从 MemBase 到 RegBase 的平滑迁移与寄存器级 VF 计算

【免费下载链接】cann-samplesCANN高性能实战演进样例与体系化调优知识库项目地址: https://gitcode.com/cann/cann-samples

本文基于 cann-samples 仓库中的KvRmsNormRopeCache实战样例,完整讲解一个融合了 RMSNorm、RoPE 与 KV Cache 写回的 AscendC 算子如何从 MemBase(以LocalTensor为中心的 UB 计算)平滑迁移到 RegBase(以RegTensor为中心的寄存器级 VF 计算)。读者将掌握 MemBase 与 RegBase 的差异边界、VF(Vector Function)设计方法、Reg::LoadAlign/Reg::StoreAlign的显式寄存器搬运范式,以及一套"外层数据流不动、热计算链下沉寄存器"的可复用调优路线。

1. 样例定位与硬件编程模型

该样例位于Samples/2_Performance/kv_rms_norm_rope_cache_story/,目录结构如下:

  • MemBase 版本:membase/full_load.asc
  • RegBase 版本:regbase/full_load.asc
  • 公共 host 侧、tiling、数据生成与精度校验:include/sample_common.h、scripts/gen_data.py
  • 构建入口:CMakeLists.txt

1.1 目标硬件与数据规格

样例面向 Ascend 950PR/950DT,编译时通过NPU_ARCH=dav-3510指定架构,固定走 AIV Vector kernel,数据类型为 BF16。典型 shape 如下表所示:

张量Shape说明
kv[B, N, S, (dv + dk)]最后一维前dv段给 RMSNorm,后dk段给 RoPE
gamma[dv]RMSNorm 缩放参数
cos/sin[B, N, S, dk]RoPE 旋转角三角函数
k_cache[B, N, S, dk]K 缓存(按 index 写回)
v_cache[B, N, S, dv]V 缓存(按 index 写回)

样例固定参数定义在 sample_common.h:SAMPLE_BATCH=8SAMPLE_NUM_HEAD=1SAMPLE_SEQ=128SAMPLE_DV=512SAMPLE_DK=128SAMPLE_UB_FACTOR=8SAMPLE_EPSILON=1e-5f、BF16 比对阈值BF16_COMPARE_TOL=6e-2f

1.2 三层存储/计算模型

从性能角度看,算子运行在三层存储/计算模型上:

  1. GM 层:输入、输出和 cache 位于全局内存,通过DataCopy/DataCopyPad与 UB 交互;
  2. UB 层:tile 数据先搬入本地 buffer,作为计算的 staging 区;
  3. Register/VF 层:RegBase 在__simd_vf__函数内使用RegTensor执行寄存器级计算。

MemBase 与 RegBase 的核心区别不在 GM/UB 搬运,而在 UB 内的计算方式:

层级MemBase 写法RegBase 写法
GM/UB 搬运DataCopyPad(GlobalTensor <-> LocalTensor)保持相同
UB stagingLocalTensorTQueTPipe保持相同
计算对象LocalTensor<T>RegTensor<T>
计算入口普通__aicore__成员函数__simd_vf__
UB 到寄存器隐含在标准 Vector API 中显式Reg::LoadAlign
寄存器到 UB隐含在标准 Vector API 中显式Reg::StoreAlign
tail 控制count、repeat、mask 参数MaskReg

因此平滑迁移的原则是:保留 MemBase 已验证过的数据流,只替换 Vector 计算部分

2. 算子计算语义:RMSNorm + RoPE + KV Cache

该算子把kv最后一维拆成两段分别处理:

kv[..., :Dv] -> RMSNorm -> v_out -> v_cache kv[..., Dv:] -> RoPE -> k_out -> k_cache

2.1 RMSNorm

RMSNorm 的计算过程为:

mean_square = mean(x * x) rms = sqrt(mean_square + epsilon) v_out = (x / rms) * gamma

2.2 RoPE

RoPE 部分把交错存储的复数拆成 real/imag 两路,再用两段 cos/sin 完成复数旋转:

real = rope_x[..., 0::2] imag = rope_x[..., 1::2] k_out_first = real * cos_first - imag * sin_first k_out_second = imag * cos_second + real * sin_second

2.3 数据生成与 golden 佐证

上述语义与 golden 生成脚本 scripts/gen_data.py 中的build_golden完全一致:rms_xkv[..., :dv]rope_xkv[..., dv:]k_out = concat(real, imag) * cos + concat(-imag, real) * sin,其中part1 = concat(real, imag)part2 = concat(-imag, real),恰好对应两段输出k_out_firstk_out_second。脚本还会根据indexk_out/v_out写回k_cache/v_cachecache_idx >= 0时写回),与内核CopyOutK/CopyOutV的 cache 写回逻辑一一对应。

当前规格下Dv=512Dk=128。RMSNorm 每行需要处理 512 个元素,包含平方、归约、开方、除法、乘gamma、BF16/FP32 转换等操作,是主要优化热点;RoPE 每行处理 128 个元素,计算链较短,但也适合用 VF 避免中间 UB 临时张量。

3. 快速编译与运行

cann-samples仓库根目录执行:

cmake -S . -B build -DNPU_ARCH=dav-3510 cmake --build build --target kv_rms_norm_rope_cache_story

构建完成后会生成两个可执行文件:

build/Samples/2_Performance/kv_rms_norm_rope_cache_story/kv_rms_norm_rope_cache_membase_full_load build/Samples/2_Performance/kv_rms_norm_rope_cache_story/kv_rms_norm_rope_cache_regbase_full_load

分别运行:

./build/Samples/2_Performance/kv_rms_norm_rope_cache_story/kv_rms_norm_rope_cache_membase_full_load ./build/Samples/2_Performance/kv_rms_norm_rope_cache_story/kv_rms_norm_rope_cache_regbase_full_load

每个可执行文件会自动生成输入数据并执行 golden 精度校验,预期最后输出包含:

PASS

从构建脚本 CMakeLists.txt 可以看到,两个.asc文件各自生成独立可执行目标,目标名规则为kv_rms_norm_rope_cache_${子目录}_${文件名},并且以-O3优化级别编译、以SOURCE_DIR宏注入源码目录,用于运行时定位gen_data.py

运行时的完整 host 流程封装在 sample_common.h 的RunSample中:aclInit→ 创建 stream → 调用GenerateDatapython3 gen_data.py --batch 8 --seq 128 --dv 512 --dk 128)→ 读取input/*.binoutput/*_golden.binaclrtMalloc+aclrtMemcpy搬入设备 → 以<<<blockNum, nullptr, stream>>>方式启动内核 → 搬回结果后对k_cachev_cachek_outv_out四路输出逐一调用CompareBf16比对,全部误差在阈值内才打印PASS

tiling 计算位于 BuildTiling:通过platform_ascendc::PlatformAscendCManager获取 AIV core 数,totalRows = B * N * SblockFactor = ceil(totalRows / coreNum)blockNum = ceil(totalRows / blockFactor)ubFactor = min(SAMPLE_UB_FACTOR, blockFactor),并将epsilonreciprocal = 1/Dv一并写入 tiling 结构体传给内核。

4. MemBase 版本:以 LocalTensor 为中心的计算链

MemBase 版本位于 membase/full_load.asc,整体采用Init -> Process -> ProcessTile的结构。

4.1 外层数据流

  • Init完成 GM 地址绑定(SetGlobalBuffer)、block 行数计算、TQue/TBuf初始化;
  • Process先预加载并转换共享gamma(BF16DataCopy后整段Cast成 FP32),随后按ubFactor循环处理 tile;
  • ProcessTile串联当前 tile 的搬入、RoPE、RMSNorm 和写回。

完整数据流如下:

Process Load gamma(BF16) -> Cast gamma(FP32) for each tile: CopyRopeAndX kv[..., :Dv] -> xLocal kv[..., Dv:] -> ropeLocal Load cos/sin Rope ropeLocal + cos/sin -> k outLocal CopyOutK k outLocal -> k_out k outLocal -> k_cache[index] RmsNorm xLocal + gammaFp32 -> v outLocal CopyOutV v outLocal -> v_out v outLocal -> v_cache[index]

在 GM/UB 搬运层,CopyRopeAndX使用两段DataCopyPad(源码):以DataCopyExtParams指定每行dk/dv个元素的跨步搬移,把后Dk维搬到ropeLocal、前Dv维搬到xLocal,从而让 RoPE 和 RMSNorm 在 UB 中独立消费各自的数据段。cossin也按 tile 搬入cosSinQueue_,与当前ropeLocal对齐。

CopyOutK/CopyOutV(源码)按行逐条DataCopyPad写回 output,并通过Mutex::Lock<PIPE_MTE3>保护对indexGm_的读取,cacheOffset >= 0时再按batch * cacheLength + cacheOffset计算 cache 地址写回。

4.2 RoPE 的 LocalTensor 实现

MemBase 的Rope(源码)仍以LocalTensor为核心:先将cos/sinCast成 FP32,再通过两次GatherMask从交错 rope 数据中分别抽取 real/imag({1, 1, NUM_EIGHT, 0}步长),随后执行两组乘加得到y0/y1Add合并后Cast回 BF16。这个流程直观、易验证,但会在cosFp32sinFp32y0y1realFp32imagFp32等 UB 临时张量之间产生多次读写。

4.3 RMSNorm 的 LocalTensor 实现

RMSNorm 是 MemBase 版本的主要热点。RmsNorm(源码)先将xLocal从 BF16 cast 到 FP32,再执行平方、分段Add归约(利用{1,1,1,NUM_EIGHT,NUM_SIXTEEN,NUM_SIXTEEN}repeat/stride 参数把 512 元素折半累加)、WholeReduceSum得到行和、乘1/Dv、加epsilonSqrtBrcb广播、逐行Div、乘gammaFp32,最后Cast回 BF16。该链路依赖多个 UB 临时区:

xFp32 -> square -> rowSum -> broadcast rms -> normalized xFp32 -> outLocal

这些临时区全部来自wsBuffer_,其分配见 Init:

pipe_.InitBuffer(wsBuffer_, ubFactor_ * (tiling_->dv * NUM_THREE + tiling_->dk * NUM_EIGHT) * sizeof(float));

即每行需要Dv*3 + Dk*8个 float 的临时空间。

MemBase 版本的仿真打点图如下:

图中展示的是一次循环的 Vector 计算内容,左边是 ROPE,右边是 RmsNorm,所有的 Vector 指令都是一个单独的 VF,整体看 Vector 指令较为细碎——这正是引入 RegBase 融合计算链的动机。

5. 从 MemBase 到 RegBase 的平滑迁移边界

从 MemBase 迁移到 RegBase 的整体原则有:

  1. 在一个完整的 VF 功能内应融尽融。理论上 VF 内融合的指令越多,UB 和寄存器交互的次数越少,性能越好。
  2. VF 内尽量少使用 Reduce 类、非对齐访存、Interleave/Deinterleave 类单发指令。
  3. VF 内尽量少使用 Membar,Membar 会导致 VF 内流水中断,使整体 IPC 不高。

迁移不是推倒重写。当前 RegBase 版本复用了 MemBase 中已经验证过的外层结构:

  • Init中的 GM tensor 绑定方式保持一致;
  • BuildTilingblockFactorubFactorblockNum保持一致;
  • CopyRopeAndX中的 GM -> UB 搬运保持DataCopyPad
  • CopyOutK/CopyOutV中的 cache 写回和 output 写回保持一致;
  • inQueuecosSinQueueoutQueue的 double buffer(BUFFER_NUM=2)数据流保持一致;
  • 输出仍走同一套 golden 校验,比较k_cachev_cachek_outv_out

真正变化的是RmsNormRope的 compute body:

MemBase: LocalTensor -> Cast/Mul/Add/WholeReduceSum/Div/Mul/Cast -> LocalTensor RegBase: __ubuf__ pointer -> Reg::LoadAlign -> RegTensor RegTensor compute chain Reg::StoreAlign -> __ubuf__ pointer

这种边界划分的好处是:

  1. 外层索引、cache offset、batch 切分和 GM 地址计算不变,降低迁移风险;
  2. 原有精度 golden 可以直接复用,MemBase 可以作为 RegBase 的对照基线;
  3. 出现错误时可以快速判断问题在 GM/UB 数据流还是 VF compute body;
  4. RegBase 的收益集中在热计算链,避免为追求寄存器化引入额外 pipeline 复杂度。

6. VF 方案设计:RegTensor 寄存器计算

RegBase 版本位于 regbase/full_load.asc。普通的RmsNorm/Rope包装函数只做一件事:通过LocalTensor.GetPhyAddr()把 UB 地址转成__ubuf__指针,再调用独立的__simd_vf__函数(RmsNorm 包装、Rope 包装)。

6.1 RMSNormVF:两段式 VF 循环

RmsNormVF(源码)把每一行的 RMSNorm 拆成两个 VF 循环。

第一段循环计算平方和:

for each row: reduceSum = 0 for i in Dv chunks: xB16 = LoadAlign(DIST_UNPACK_B16) xFp32 = Cast(BF16 -> FP32) square = xFp32 * xFp32 chunk = ReduceSum(square) reduceSum += chunk

第二段循环归一化并乘gamma

rms = sqrt(reduceSum / Dv + epsilon) for i in Dv chunks: xB16 = LoadAlign(DIST_UNPACK_B16) gammaB16 = LoadAlign(DIST_UNPACK_B16) xFp32 = Cast(BF16 -> FP32) gammaFp32 = Cast(BF16 -> FP32) norm = (xFp32 / rms) * gammaFp32 outB16 = Cast(FP32 -> BF16) StoreAlign(DIST_PACK_B32)

对应到源码,chunk 粒度由VL_FP32_SIZE = 256 / sizeof(float) = 64决定,dvLoop_ = (dv + VL_FP32_SIZE - 1) / VL_FP32_SIZE,即Dv=512时每行 8 个 chunk。第一段用ReduceSum(reduceSum, reduceSum, fullMask)把 64-lane 的行内和归约到单一值;随后MulsreciprocalAddsepsilonSqrt开方、Div(invRms, one, rmsValue)取倒数,再用Duplicate<HighLowPart::LOWEST>把标量广播成invRmsBrc供整行乘除使用;第二段按 chunk 同时LoadAlignx 与 gamma,Mul(x, invRmsBrc)Mul(norm, gammaFp32)两级乘后Cast回 BF16 并StoreAlign

设计要点:

  • BF16 输入先升到 FP32 做平方、归约、除法和乘法,保证 RMSNorm 中间计算精度;
  • reduceSumrmsValueinvRmsinvRmsBrcnorm等短生命周期中间量全部保留在RegTensor中;
  • gamma不再预先 cast 成一整段 FP32 UB 临时张量,而是在 VF 中按 chunk load BF16 并 cast 到 FP32;
  • 通过UpdateMask<float>(remaining)控制最后一个 chunk 的有效 lane,避免把 tail padding 当有效数据。

6.2 RopeVF:寄存器内复数旋转

RopeVF(源码)的核心思路是直接从 BF16 rope 数据中拆出 real/imag,并用两段 cos/sin 完成复数旋转:

real, imag = LoadAlign(DIST_DINTLV_B16, rope) realFp32 = Cast + Interleave imagFp32 = Cast + Interleave out_first = realFp32 * cos_first - imagFp32 * sin_first out_second = imagFp32 * cos_second + realFp32 * sin_second StoreAlign(out_first) StoreAlign(out_second)

从源码实现看,real/imag 的拆分是通过两次带不同RegLayout的 Cast 完成的:CAST_B16_TO_B32RegLayout::ZERO,取 even lane)得到realFp32CAST_B16_TO_B32_ODDRegLayout::ONE,取 odd lane)得到imagFp32(cast trait 定义),效果等价于对交错数据做 deinterleave。随后Mul/Sub/Add全部在寄存器内串联,两次StoreAlign分别写回outFirstoutSecond

设计要点:

  • real/imag、cos/sin、mul/result 都在寄存器内串联,无 UB 中间张量;
  • RoPE 的pairCount = Dk / 2 = 64mask = UpdateMask<float>(pairCount)固定控制 64 个 FP32 lane;
  • 输出分成前后两段写回(包装函数中以out + pairCount作为第二段起始地址),与 golden 中concat(real, imag)的布局保持一致。

7. 理论分析与实测验证

7.1 UB 使用量减少

MemBase RMSNorm 需要多个 UB 临时张量:

xFp32 -> square -> rowSum -> broadcast rms -> xFp32 norm -> outLocal

并且标准 Vector API 之间通常会形成 UB load/store 边界。RegBase 把平方、归约、除法、乘 gamma 等中间状态尽量保存在寄存器中,只在 VF 入口 load 输入、在 VF 末尾 store 输出。

理论收益来自:

  • 减少中间结果落 UB;
  • 减少多段 Vector API 之间的 UB 往返;
  • 降低 UB 临时 buffer 占用。

对比两个版本的 buffer 分配:

MemBase: gammaQueue: BF16 gamma + FP32 gamma wsBuffer_: rows * (Dv * 3 + Dk * 8) * sizeof(float) RegBase: gammaQueue: BF16 gamma no wsBuffer_

RegBase 的 Init 中已无TBuf<TPosition::VECCALC> wsBuffer_成员,这是最直接的 UB 占用优化。

7.2 精度路径保持一致

样例输入和输出是 BF16,但 RMSNorm 的关键中间计算必须使用 FP32:

BF16 load -> FP32 compute -> BF16 store

RegBase 版本通过CAST_B16_TO_B32CAST_B32_TO_B16显式表达转换路径,其中回写 cast 使用RoundMode::CAST_RINT(round to nearest)。实测中,MemBase 和 RegBase 均通过同一套 golden 校验:

MemBase: k_cache/v_cache/k_out/v_out 全部 PASS RegBase: k_cache/v_cache/k_out/v_out 全部 PASS

其中 RegBasev_cache/v_out的最大 diff 为0.0078125,远低于 sample_common.h 中定义的 BF16 比对阈值6e-2

8. 实际调优路径

8.1 保持 full-load 外层数据流

当前样例一次 tile 处理ubFactor=8行,运行时实测打印:

blockNum 54, blockFactor 19, ubFactor 8

这与 BuildTiling 的推导吻合:总行数8 * 1 * 128 = 1024,按 AIV core 数向上取整得到每个 block 约 19 行,UB 内一次搬入最多 8 行。这个切分对 MemBase 和 RegBase 共同适用,因此调优时先不改外层,只验证 compute body 的收益和正确性。

8.2 减少 UB 使用

MemBase 版本需要wsBuffer_存放 RMSNorm 和 RoPE 的 FP32 中间结果;RegBase 版本删除了这个大块 VECCALC,改为在 VF 中使用RegTensor临时变量。对比:

MemBase: gammaQueue: BF16 gamma + FP32 gamma wsBuffer_: rows * (Dv * 3 + Dk * 8) * sizeof(float) RegBase: gammaQueue: BF16 gamma no wsBuffer_

这是最直接的 UB 占用优化,为 double buffer 深度或更大的ubFactor留出了容量空间。

8.3 gamma 处理:整段预转换改为 VF 内按需转换

MemBase 版本在主循环前把整段gamma从 BF16 cast 成 FP32,并保存在 UB 中供后续每行复用;RegBase 版本保留 BF16gammaLocal,在RmsNormVF的第二段循环中按 chunk load 并 cast。这样做的取舍:

  • 优点:减少一份完整 FP32 gamma UB buffer;
  • 优点:gamma 与 x 的 load/cast 粒度一致,寄存器链更紧凑;
  • 代价:每行都会重新 cast gamma chunk。

当前样例的主要目标是演示 RegBase VF 迁移和减少 UB 中间张量;如果后续 profiler 证明 gamma cast 成本突出,可以考虑引入小块 gamma 预处理或更细粒度复用策略,但不能破坏 UB 容量和 double buffer 数据流。

8.4 切换验证顺序

建议按以下顺序调优,避免一次修改多个变量后无法归因:

  1. 固定输入数据和 shape,只替换 compute body,确认精度 PASS;
  2. 对比 MemBase 和 RegBase 的输出 diff,先确认 BF16 阈值内一致;
  3. 使用 profiler 观察 UB load/store、Vector 指令、pipeline stall;
  4. 若 UB 压力仍高,继续检查 VF 内是否还有可消除的 store/load;
  5. 若算术指令成为瓶颈,再检查 reduce 组织、gamma cast 复用和 div/sqrt 链路;
  6. 若搬运成为瓶颈,再调整ubFactor、double buffer 深度或 GM copy 形态。

当前已完成的功能验证:

kv_rms_norm_rope_cache_membase_full_load: PASS kv_rms_norm_rope_cache_regbase_full_load: PASS

9. RegBase 版本运行效果

RegBase 版本位于 regbase/full_load.asc。仿真打点图中左侧是 rope,已经融合成 1 个 VF;右侧是 RmsNorm,融合成 1 个 VF。

选择 Rope 中 RVEC_EX 部分的指令,从下侧的简易统计可以估算出 IPC 在 1.36 左右:

选择 RmsNorm 中 RVEC_EX 部分的指令,估算出 IPC 在 1.15 左右。RmsNorm 的 IPC 比 Rope 低,是因为其中有单发指令,而且代码中第一个循环体前后有依赖,所以整体 IPC 比 Rope 低:

需要说明的是,上述 IPC 数值来自仿真打点图的简易统计估算,用于对比两个 VF 内部的指令密度差异;实际性能收益应结合 profiler 数据进一步确认。

10. 迁移实践清单

从 MemBase 切到 RegBase 时,可以按以下清单执行:

  1. 先保留 MemBase 的 host、tiling、queue、CopyIn、CopyOut;
  2. 找出最热的LocalTensor计算链,本样例是 RMSNorm,其次是 RoPE;
  3. 把计算链拆成独立__simd_vf__函数;
  4. 在普通__aicore__函数中只做LocalTensor.GetPhyAddr()__ubuf__指针的转换;
  5. VF 内只使用RegTensorMaskRegReg::Load*Reg::Store*Reg::*compute API;
  6. BF16 路径显式写清BF16 -> FP32 -> BF16的 cast trait 和 round mode;
  7. tail 只用MaskReg控制,不依赖 UB padding 参与数学;
  8. 每次优化后都运行同一套 golden 校验,确认k_cachev_cachek_outv_out全部 PASS。

一句话总结:RegBase 的平滑迁移路线是"外层不动、热链下沉、寄存器串联、同源校验"。本样例正是沿着这条路线,把 MemBase 中以LocalTensor为中心的计算链切换为 VF 中以RegTensor为中心的寄存器计算链——先通过同一套 golden 校验守住精度底线,再借助仿真打点观察 VF 融合效果,最终在减少 UB 中间张量的同时让热计算链在寄存器内连续执行。

【免费下载链接】cann-samplesCANN高性能实战演进样例与体系化调优知识库项目地址: https://gitcode.com/cann/cann-samples

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询