CANN ops-transformer 算子实战:mamba2_chunk_cumsum 在 MambaV2 Prefill 阶段的分块累积求和实现解析
2026/9/18 20:55:17 网站建设 项目流程

CANN ops-transformer 算子实战:mamba2_chunk_cumsum 在 MambaV2 Prefill 阶段的分块累积求和实现解析

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

mamba2_chunk_cumsum 是 CANN ops-transformer 仓库experimental/mamba目录下 MambaV2 系列定制算子之一,面向华为昇腾 NPU(910B 平台)实现。本算子位于 MambaV2 Prefill 阶段 chunk 计算流水线的起点:它按 chunk 对输入序列执行因果顺序的累积求和(cumulative sum),产出时间步衰减量dtout、累积状态量dacs以及每个 chunk 的末状态dacs_chunk,为后续 chunk 状态更新(chunk_state)与 selective scan(chunk_scan)提供输入。读完本文,你将掌握该算子的数学语义、I/O 形状与数据类型、PyTorch 调用方式、Vector 流水实现原理以及精度/性能验证方法。

功能定位:MambaV2 Prefill 分块流水线的起点

Mamba 系列基于状态空间模型(SSM)以线性复杂度 $O(N)$ 替代 Transformer 的 $O(N^2)$ 自注意力,Mamba v2 进一步通过状态空间对偶性(SSD)将递推计算改写为结构化矩阵乘法,从而支持分块并行。experimental/mamba目录下的 mamba2_chunk_xxx 四个算子正是 Prefill 阶段 chunk 计算的核心实现模块,对应关系如下(详见 experimental/mamba/Readme.md):

本目录算子vLLM 对应模块功能
mamba2_chunk_cumsumssd_combined(cumsum 部分)chunk 内累积求和,用于状态递推
mamba2_chunk_statessd_chunk_statechunk 内离散状态更新
mamba2_chunk_state_passingssd_state_passing跨 chunk 状态传递与衰减
mamba2_chunk_scanssd_combined(scan 部分)selective scan 扫描,结合状态与门控

mamba2_chunk_cumsum 是这条流水线的第一步。MambaV2 的 chunked 计算策略将长序列按chunk_size拆分为若干 chunk:chunk 内部(Intra-chunk)展开为可并行计算的形式,chunk 之间(Inter-chunk)通过递推状态传递保持 SSM 的线性特性。本算子负责在 chunk 内部沿时间步方向做因果累积求和,将"每个时间步的衰减增量"累加成"截至该时间步的累积衰减量",供后续 chunk 状态更新与 selective scan 使用。原文档明确指出:算子对输入序列在 S 维度按 chunk_size 拆分,并在每个 chunk 内按照因果顺序执行 cumulative sum。

计算语义:从测试参考实现看算子数学公式

原 README 未给出逐行公式,但仓库中的 PyTorch 参考实现完整揭示了算子的数学语义。tests/test_chunk_cumsum.py中的mamba2_chunk_cumsum_forward是官方提供的精度比对基准(test_chunk_cumsum.py):

def mamba2_chunk_cumsum_forward(at, dt, dtbias, dtmask): B, C, L, H = dt.shape dt_add = dt.to(torch.float32) + torch.reshape(dtbias.to(torch.float32), (1, 1, 1, H)) dtout = torch.log(1 + torch.exp(dt_add)) # softplus dtout = torch.where(dt_add < 20, dtout, dt_add) # 大值线性近似,数值稳定 dtout = torch.clamp(dtout, max=10000000.0) * dtmask.to(torch.float32) da = dtout * torch.reshape(at, (1, 1, 1, H)) # 与 at 逐元素相乘 dacs = torch.cumsum(da, dim=2) # 沿 L(时间步)维度累积 dacs_chunk = torch.reshape(dacs[:, :, -1, :], (B, C, 1, H)) return dtout, dacs, dacs_chunk

由此可以总结出算子的完整计算链:

  1. 加偏置dt(FP16,转 FP32 后)与dt_bias(形状(H,),广播到(1,1,1,H))相加得到dt_add
  2. Softplus 激活:对dt_add计算log(1 + exp(dt_add))得到步长正数化结果;当dt_add >= 20时直接取dt_add(softplus 的线性近似),这是避免exp溢出的数值稳定处理,与 Kernel 中常量COMPARE_VALUE = 20.0一一对应;
  3. Clamp 与掩码:将结果裁剪到上限10000000.0(对应 Kernel 常量CLAMP_MAX),再乘以dt_mask
  4. 加权dtoutat逐元素相乘,得到每个时间步的增量da
  5. 因果累积求和:沿 L(chunk 内时间步)维度执行torch.cumsum(da, dim=2),得到dacs
  6. chunk 末状态:取每个 chunk 最后一个时间步的累积值dacs[:, :, -1, :],reshape 为(B, C, 1, H)dacs_chunk,它是 chunk 内状态递推的最终结果,将作为下一阶段跨 chunk 状态传递(chunk_state_passing)的输入之一。

需要说明的是,以上公式由测试脚本中的参考实现推断得出,属于对算子数学语义的源码级还原。

Kernel 输入输出(I/O)

原 README 给出的输入输出规格如下:

输入

Tensorshapedtype
atHFP32
dtBCLHFP16
dt_biasHFP16
dt_maskBCLHFP16

输出

Tensorshapedtype
dtoutBCLHFP32
dacsBCLHFP32
dacs_chunkBCHFP32

在 torch_interface.cpp 中可以看到输出张量的实际分配方式:dt_outdacs_out均为{B, C, L, H}的 FP32 空张量,dacs_chunk_out{B, C, 1, H}的 FP32 空张量(注意 README 表格中写为BCH,实际实现将第三个维度保留为 1,形状为(B, C, 1, H))。在进入 Kernel 前,入口函数还会对输入做一步类型规整:at转为 FP32、dt/dt_bias/dt_mask转为 FP16(torch_interface.cpp),注释说明"可能并非必需",属于防御性转换。

参数说明

参数含义
Bbatch size
Cnumber of chunks(chunk 数量)
Lchunk size(每个 chunk 内的时间步数)
Hnumber of head(通道/头数)

其中C*L 为 padding 后的序列长度:MambaV2 要求序列长度是 chunk_size 的整数倍,调用前需要在 S 维度对序列做 padding。从experimental/mamba/Readme.md的特性说明可知,当前版本所有 mamba2_chunk_xxx 系列算子均支持 BSND 数据布局,S 维度需在调用前 pad 至 chunk_size 的整数倍,且当前版本仅支持固定 chunk_size = 256(即测试用例中L = 256)。

调用方式

在安装了npu_ops_transformer_exttorch 扩展后,可按原文档方式调用:

import npu_ops_transformer_ext out = torch.ops.npu_ops_transformer_ext.mamba2_chunk_cumsum(at, dt, dt_bias, dt_mask)

需要提醒的是:README 中记录的算子名为mamba2_chunk_cumsum,而实际 Kernel 注册名与测试脚本使用的是mambav2_chunk_cumsum(见 torch_interface.cpp 的m.impl("mambav2_chunk_cumsum", ...)以及 test_chunk_cumsum.py 的调用torch.ops.npu_ops_transformer_ext.mambav2_chunk_cumsum(...)),CMakeLists 中定义的算子名同为mambav2_chunk_cumsum(CMakeLists.txt)。如果按 README 中的名字调用报未找到算子,请改用注册名mambav2_chunk_cumsum

算子的完整注册包含两套实现:PrivateUse1设备上的 NPU 实现和Meta上的 meta 函数(用于 shape 推导),meta 函数仅做输入有效性检查并原样返回(torch_interface.cpp)。

源码级实现解析:Vector 流水如何完成累积求和

原文档指出"本算子基于 Vector 实现累积求和计算"。结合 op_kernel/CustVec.h 与 torch_interface.cpp,可以从三个层面还原其实现:

1. 启动与资源准备(Host 侧)

mambav2_chunk_cumsum入口函数完成以下工作:

  • 解析dt的四个维度得到B, C, L, H
  • 20 个 block(blockDims = 20启动 Kernel(torch_interface.cpp),算子通过GetBlockNum()参与核内 tiling 计算;
  • 分配 1024 字节用户 workspace 加上平台系统 workspace(GetLibApiWorkSpaceSize()),合并为workspaceTensor传入 Kernel(torch_interface.cpp);
  • 通过at_npu::native::OpCommand::RunOpApiV2("Mambav2ChunkCumsum", acl_call)提交内核执行(torch_interface.cpp)。

2. Kernel 内分核与 tiling(AI Core 侧)

CustVec.h中定义了分块常量与 tiling 逻辑:

  • 基础分块常量:BASEL = 64(L 维每次处理 64 行)、BASEH = 128SUB_BASEH = 64TILE_BLK_SIZE = BASEL * SUB_BASEH(CustVec.h);
  • 数值常量:COMPARE_VALUE = 20.0(softplus 线性近似阈值)、CLAMP_MAX = 10000000.0f(clamp 上限),与前述测试参考实现完全一致(CustVec.h);
  • tilingShapeCustVec按 H 大小自适应分块(CustVec.h):H ≤ 32 时effective_H = 32nstepsH = 1;H ≤ 64 时effective_H = 64;H ≤ 128 时拆成 2 个 64 的子块;更大 H 则按CeilDiv(H, 128)分步。每个 AI Core 负责BCH = B * C * nstepsH中均分的一段(BCH_PER_CORE = CeilDiv(BCH, GetBlockNum())),并以双缓冲(DBuff)方式流水执行。

3. Vector 计算流水(Process_dacs)

Compute()对每个(b, c, h)组合按BASEL分块循环,Process_dacs中依次执行(CustVec.h):

  1. Castdtdt_biasdt_mask从 FP16 转 FP32;
  2. Add累加dt_bias(通过src1RepStride = 0实现标量广播);
  3. CompareScalar+Exp+Adds(1.0f)+Ln计算 softplus,再用Select依据dt_add < 20选择 softplus 值或线性近似值;
  4. 再次CompareScalar+Select完成CLAMP_MAX截断;
  5. Mul依次乘dt_mask、乘at
  6. 累积求和:借助cumsum_tensor(chunk 级累加器)跨 BASEL 块传递部分和——块内对 63 个偏移执行逐行Add(每个时间步累加上一行),块末将最后一行保存回cumsum_tensor,实现因果累积;
  7. 数据搬运上通过DEvent<PIPE_MTE2, PIPE_V>DEvent<PIPE_V, PIPE_MTE3>等事件同步 MTE2(GM→UB 搬运)、Vector 计算与 MTE3(UB→GM 写回),并以双缓冲隐藏搬运延迟(CustVec.h)。

dacs_chunk(每个 chunk 末时间步的累积值)在Move_ub2gm中当(l + BASEL) >= L时单独写回out2_mtx(CustVec.h)。

4. 公共工具:张量搬运与数据类型

CustVec.h依赖的GM2UB/UB2GM/UB2UB等搬运封装与CeilDiv、双缓冲DBuff、事件DEvent等基础设施位于公共头文件 common/tensorutils.h 与 common/paramutils.h,后者提供 FP16/FP32 互转及默认 unary/binary 的 repeat 参数(STRIDE_FLOAT = 8CAST_STRIDE_HALF = 4),供本算子 Kernel 直接复用。

测试与精度验证

算子测试脚本位于 tests/test_chunk_cumsum.py,运行方式与原文档一致:

python test_chunk_cumsum.py

脚本以B=1, C=4, H=128, G=8, L=256的典型配置(chunk_size 固定为 256)构造随机输入,将 NPU Kernel 输出与上述 PyTorch 参考实现逐一比对:

dtout, dacs, dacs_chunk = mamba2_chunk_cumsum_forward(tensor_at, tensor_dt, tensor_dtbias, tensor_dtmask) npu_dtout, npu_dacs, npu_dacs_chunk = torch.ops.npu_ops_transformer_ext.mambav2_chunk_cumsum(...) check_diff(dtout.cpu(), npu_dtout.cpu()) check_diff(dacs.cpu(), npu_dacs.cpu()) check_diff(dacs_chunk.cpu(), npu_dacs_chunk.cpu())

check_diffprofiling工具来自 utils/utils.py:前者打印最大绝对误差与相对误差;后者通过torch_npu.profiler对 Torch 参考实现与 NPU Kernel 分别做 5 次 warmup、10 次计时,输出单次平均耗时(微秒),并落盘 profile 结果到TORCH_profile_results/NPU_KERNEL_profile_results目录,用于精度与性能的双重验证。

编译与运行环境

mamba2_chunk_cumsum作为 torch 扩展算子编译,CMakeLists.txtBUILD_TORCH_OPS开关打开时构建,算子目标名为mambav2_chunk_cumsum,源文件以--npu-arch=dav-2201(昇腾 910B 系列 AI Core 架构)并链接 tiling_api/platform/register 库进行编译(CMakeLists.txt)。整仓编译与安装流程见 experimental/mamba/Readme.md:

cd experimental/npu_ops_transformer_ext python3 -m build --wheel -n cd dist pip3 install *.whl --force-reinstall --no-deps

编译安装后即可在测试目录中运行python test_chunk_cumsum.py

使用注意事项

  • 序列长度对齐:输入dt/dt_mask的 S 维度(即C*L)必须是 chunk_size 的整数倍,调用前需在 S 维度 padding;
  • chunk_size 固定:当前版本仅支持固定 chunk_size = 256,测试用例也以L=256验证;
  • 精度支持:算子支持 FP16/FP32 输入输出——at为 FP32,dt/dt_bias/dt_mask为 FP16,三个输出均为 FP32(中间计算在 FP32 下进行,参考实现同样先将 FP16 输入升到 FP32);
  • 算子命名:README 中的mamba2_chunk_cumsum与实现/测试使用的注册名mambav2_chunk_cumsum存在差异,实际调用以注册名为准;
  • 在整条流水线中的位置:本算子输出dacs/dacs_chunk/dtout,其中dacsdtout会被 mamba2_chunk_state 用于 chunk 内状态递推,dacs还会在 mamba2_chunk_state_passing 中用于跨 chunk 状态衰减,而 mamba2_chunk_scan 则结合dacs/dtout与门控生成 chunk 有效输出。理解这一上下游关系,有助于把握该算子在 MambaV2 Prefill 全流程中的具体贡献。

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

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

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

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

立即咨询