十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

CANN ops-transformer 算子实战:mamba2_chunk_cumsum 在 MambaV2 Prefill 阶段的分块累积求和实现解析 CANN ops-transformer 算子实战mamba2_chunk_cumsum 在 MambaV2 Prefill 阶段的分块累积求和实现解析【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformermamba2_chunk_cumsum 是 CANN ops-transformer 仓库experimental/mamba目录下 MambaV2 系列定制算子之一面向华为昇腾 NPU910B 平台实现。本算子位于 MambaV2 Prefill 阶段 chunk 计算流水线的起点它按 chunk 对输入序列执行因果顺序的累积求和cumulative sum产出时间步衰减量dtout、累积状态量dacs以及每个 chunk 的末状态dacs_chunk为后续 chunk 状态更新chunk_state与 selective scanchunk_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_combinedcumsum 部分chunk 内累积求和用于状态递推mamba2_chunk_statessd_chunk_statechunk 内离散状态更新mamba2_chunk_state_passingssd_state_passing跨 chunk 状态传递与衰减mamba2_chunk_scanssd_combinedscan 部分selective scan 扫描结合状态与门控mamba2_chunk_cumsum 是这条流水线的第一步。MambaV2 的 chunked 计算策略将长序列按chunk_size拆分为若干 chunkchunk 内部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.pydef 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, max10000000.0) * dtmask.to(torch.float32) da dtout * torch.reshape(at, (1, 1, 1, H)) # 与 at 逐元素相乘 dacs torch.cumsum(da, dim2) # 沿 L时间步维度累积 dacs_chunk torch.reshape(dacs[:, :, -1, :], (B, C, 1, H)) return dtout, dacs, dacs_chunk由此可以总结出算子的完整计算链加偏置dtFP16转 FP32 后与dt_bias形状(H,)广播到(1,1,1,H)相加得到dt_addSoftplus 激活对dt_add计算log(1 exp(dt_add))得到步长正数化结果当dt_add 20时直接取dt_addsoftplus 的线性近似这是避免exp溢出的数值稳定处理与 Kernel 中常量COMPARE_VALUE 20.0一一对应Clamp 与掩码将结果裁剪到上限10000000.0对应 Kernel 常量CLAMP_MAX再乘以dt_mask加权dtout与at逐元素相乘得到每个时间步的增量da因果累积求和沿 Lchunk 内时间步维度执行torch.cumsum(da, dim2)得到dacschunk 末状态取每个 chunk 最后一个时间步的累积值dacs[:, :, -1, :]reshape 为(B, C, 1, H)即dacs_chunk它是 chunk 内状态递推的最终结果将作为下一阶段跨 chunk 状态传递chunk_state_passing的输入之一。需要说明的是以上公式由测试脚本中的参考实现推断得出属于对算子数学语义的源码级还原。Kernel 输入输出I/O原 README 给出的输入输出规格如下输入TensorshapedtypeatHFP32dtBCLHFP16dt_biasHFP16dt_maskBCLHFP16输出TensorshapedtypedtoutBCLHFP32dacsBCLHFP32dacs_chunkBCHFP32在 torch_interface.cpp 中可以看到输出张量的实际分配方式dt_out与dacs_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转为 FP16torch_interface.cpp注释说明可能并非必需属于防御性转换。参数说明参数含义Bbatch sizeCnumber of chunkschunk 数量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_cumsumCMakeLists.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 个 blockblockDims 20启动 Kerneltorch_interface.cpp算子通过GetBlockNum()参与核内 tiling 计算分配 1024 字节用户 workspace 加上平台系统 workspaceGetLibApiWorkSpaceSize()合并为workspaceTensor传入 Kerneltorch_interface.cpp通过at_npu::native::OpCommand::RunOpApiV2(Mambav2ChunkCumsum, acl_call)提交内核执行torch_interface.cpp。2. Kernel 内分核与 tilingAI Core 侧CustVec.h中定义了分块常量与 tiling 逻辑基础分块常量BASEL 64L 维每次处理 64 行、BASEH 128、SUB_BASEH 64、TILE_BLK_SIZE BASEL * SUB_BASEHCustVec.h数值常量COMPARE_VALUE 20.0softplus 线性近似阈值、CLAMP_MAX 10000000.0fclamp 上限与前述测试参考实现完全一致CustVec.htilingShapeCustVec按 H 大小自适应分块CustVec.hH ≤ 32 时effective_H 32且nstepsH 1H ≤ 64 时effective_H 64H ≤ 128 时拆成 2 个 64 的子块更大 H 则按CeilDiv(H, 128)分步。每个 AI Core 负责BCH B * C * nstepsH中均分的一段BCH_PER_CORE CeilDiv(BCH, GetBlockNum())并以双缓冲DBuff方式流水执行。3. Vector 计算流水Process_dacsCompute()对每个(b, c, h)组合按BASEL分块循环Process_dacs中依次执行CustVec.hCast将dt、dt_bias、dt_mask从 FP16 转 FP32Add累加dt_bias通过src1RepStride 0实现标量广播CompareScalarExpAdds(1.0f)Ln计算 softplus再用Select依据dt_add 20选择 softplus 值或线性近似值再次CompareScalarSelect完成CLAMP_MAX截断Mul依次乘dt_mask、乘at累积求和借助cumsum_tensorchunk 级累加器跨 BASEL 块传递部分和——块内对 63 个偏移执行逐行Add每个时间步累加上一行块末将最后一行保存回cumsum_tensor实现因果累积数据搬运上通过DEventPIPE_MTE2, PIPE_V、DEventPIPE_V, PIPE_MTE3等事件同步 MTE2GM→UB 搬运、Vector 计算与 MTE3UB→GM 写回并以双缓冲隐藏搬运延迟CustVec.h。dacs_chunk每个 chunk 末时间步的累积值在Move_ub2gm中当(l BASEL) L时单独写回out2_mtxCustVec.h。4. 公共工具张量搬运与数据类型CustVec.h依赖的GM2UB/UB2GM/UB2UB等搬运封装与CeilDiv、双缓冲DBuff、事件DEvent等基础设施位于公共头文件 common/tensorutils.h 与 common/paramutils.h后者提供 FP16/FP32 互转及默认 unary/binary 的 repeat 参数STRIDE_FLOAT 8、CAST_STRIDE_HALF 4供本算子 Kernel 直接复用。测试与精度验证算子测试脚本位于 tests/test_chunk_cumsum.py运行方式与原文档一致python test_chunk_cumsum.py脚本以B1, C4, H128, G8, L256的典型配置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_diff与profiling工具来自 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.txt在BUILD_TORCH_OPS开关打开时构建算子目标名为mambav2_chunk_cumsum源文件以--npu-archdav-2201昇腾 910B 系列 AI Core 架构并链接 tiling_api/platform/register 库进行编译CMakeLists.txt。整仓编译与安装流程见 experimental/mamba/Readme.mdcd 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 维度 paddingchunk_size 固定当前版本仅支持固定 chunk_size 256测试用例也以L256验证精度支持算子支持 FP16/FP32 输入输出——at为 FP32dt/dt_bias/dt_mask为 FP16三个输出均为 FP32中间计算在 FP32 下进行参考实现同样先将 FP16 输入升到 FP32算子命名README 中的mamba2_chunk_cumsum与实现/测试使用的注册名mambav2_chunk_cumsum存在差异实际调用以注册名为准在整条流水线中的位置本算子输出dacs/dacs_chunk/dtout其中dacs、dtout会被 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),仅供参考
返回列表