CANN ops-math RightShift 算子 aclnn 接口使用指南:两段式调用、参数约束与源码级原理
算子库人工智能CANN【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-math点击查看免费下载本文以开源仓库 cann/ops-math 中 experimental/math/right_shift/docs/aclnnRightShift.md 为核心依据系统讲解数学基础算子库中 RightShift按位右移算子的功能语义、两段式 aclnn 接口的完整参数与返回码、约束条件及可编译运行的调用示例并结合仓库中的 op_api、op_host、op_kernel 三层源码深入解析参数校验、类型推导、broadcast 处理与 tiling 切分的底层实现。读完本文读者可以独立完成 RightShift 算子在 Atlas A2 系列产品上的工程化调用并理解其内部计算流程。功能说明与数学语义RightShift 算子对输入张量input中的每个元素按照shiftBits中对应位置的移位位数执行按位右移计算公式为$$ out_i input_i \gg shiftBits_i $$其中out为输出张量input为被右移的输入张量shiftBits为右移位数张量。该语义与 experimental/math/right_shift/README.md 中的描述完全一致README 中记作z_i x_i y_i。算子行为遵循三条核心规则broadcast 语义input与shiftBits支持广播broadcast输出out的 shape 为二者 broadcast 后的 shape。移位方向语义对有符号整数执行算术右移符号位扩展对无符号整数执行逻辑右移高位补零。非法移位位数规则当移位位数不在合法范围内时不直接产生未定义行为而是输出确定值有符号整数若shiftBits_i 0或shiftBits_i bitWidthbitWidth为该整数类型的位宽则当input_i 0时输出-1否则输出0无符号整数若shiftBits_i bitWidth则输出0。该规则的实现可以在 kernel 源码中直接找到证据experimental/math/right_shift/op_kernel/right_shift.h 中的IsInvalidShiftshiftValue sizeof(T) * 8 - 1或有符号类型下shiftValue 0判定为非法与InvalidShiftValue有符号时xValue 0 ? -1 : 0无符号恒为0两个函数与实际文档语义逐条对应。产品支持情况产品是否支持Atlas A2 训练系列产品/Atlas 800I A2 推理产品/A200I A2 Box 异构组件√从算子注册代码看该算子的 AI Core 配置针对ascend910b平台注册见 op_host/right_shift_def.cpp 中AddConfig(ascend910b, aicoreConfig)与文档所列产品族一致UT 测试也通过op::SetPlatformSocVersion(op::SocVersion::ASCEND910B)指定平台见 tests/ut/op_api/test_aclnn_right_shift.cpp。两段式接口与函数原型RightShift 算子对外提供的是 CANN 标准的两段式 aclnn 接口详见仓库文档 docs/zh/context/two_phase_api.md第一段aclnnRightShiftGetWorkspaceSize完成入参校验并获取计算所需 workspace 大小以及包含算子计算流程的执行器executor第二段aclnnRightShift使用第一段返回的 workspace 与 executor在指定 Stream 上执行计算。aclnnStatus aclnnRightShiftGetWorkspaceSize( const aclTensor *input, const aclTensor *shiftBits, aclTensor *out, uint64_t *workspaceSize, aclOpExecutor **executor)aclnnStatus aclnnRightShift( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)两个接口的声明位于 op_api/aclnn_right_shift.h实现在 op_api/aclnn_right_shift.cpp。aclnnRightShiftGetWorkspaceSize 参数说明参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续Tensorinput输入待右移的输入张量对应公式中的input需要与shiftBits满足 broadcast 关系INT8、UINT8、INT16、UINT16、INT32、UINT32、INT64、UINT64ND0-8√shiftBits输入右移位数张量对应公式中的shiftBits需要与input满足 broadcast 关系INT8、UINT8、INT16、UINT16、INT32、UINT32、INT64、UINT64ND0-8√out输出右移计算结果对应公式中的outshape 需要与input和shiftBitsbroadcast 后的 shape 一致INT8、UINT8、INT16、UINT16、INT32、UINT32、INT64、UINT64ND0-8√workspaceSize输出返回需要在 Device 侧申请的 workspace 大小-----executor输出返回 op 执行器包含算子计算流程-----第一段接口会完成入参校验校验逻辑在源码中体现为CheckParams组合了CheckNotNull、CheckDtypeValid、CheckShape三个检查见 op_api/aclnn_right_shift.cpp出现以下场景时报错返回码错误码描述ACLNN_ERR_PARAM_NULLPTR161001input、shiftBits、out、workspaceSize或executor为空指针ACLNN_ERR_PARAM_INVALID161002input、shiftBits或out的数据类型不在支持范围内ACLNN_ERR_PARAM_INVALID161002input或shiftBits的维度超过 8 维ACLNN_ERR_PARAM_INVALID161002input与shiftBits不满足 broadcast 关系ACLNN_ERR_PARAM_INVALID161002out的 shape 与 broadcast 后的 shape 不一致ACLNN_ERR_PARAM_INVALID161002input与shiftBits无法完成类型推导或推导后的类型无法转换为out的数据类型ACLNN_ERR_INNER_CREATE_EXECUTOR-创建执行器失败ACLNN_ERR_INNER_NULLPTR-内部计算流程创建失败返回码的完整说明可参考仓库文档 docs/zh/context/aclnn_return_code.md。从源码看MAX_INPUT_DIM被定义为 8DTYPE_SUPPORT_LIST精确列出上述 8 种整数类型CheckShape内部通过OP_CHECK_BROADCAST_AND_INFER_SHAPE完成 broadcast 关系校验并将推导出的 shape 与out逐维比对op_api/aclnn_right_shift.cpp。aclnnRightShift 参数说明参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址workspaceSize输入在 Device 侧申请的 workspace 大小由第一段接口aclnnRightShiftGetWorkspaceSize获取executor输入op 执行器包含算子计算流程stream输入指定执行任务的 Stream第二段接口的实现在 op_api/aclnn_right_shift.cpp先校验executor非空再通过CommonOpExecutorRun(workspace, workspaceSize, executor, stream)统一执行器运行接口把计算任务下发到指定 Stream。约束说明仅支持整数类型INT8、UINT8、INT16、UINT16、INT32、UINT32、INT64、UINT64。input、shiftBits、out均支持 0-8 维 Tensor0 维表示标量。input与shiftBits需要满足 broadcast 关系out的 shape 需要等于 broadcast 后的 shape。支持空 Tensor元素个数为 0 时跳过计算并返回空结果。对应源码中input-IsEmpty() || shiftBits-IsEmpty()时直接置*workspaceSize 0并返回成功op_api/aclnn_right_shift.cpp。支持非连续 Tensor接口内部会进行连续化处理。这一点在约束说明、READMEexperimental/math/right_shift/README.md以及通用文档 docs/zh/context/non_contiguous_tensor.md 中均有体现。接口支持input与shiftBits进行类型推导promote推导后的数据类型需能转换为out的数据类型。调用示例下面为文档给出的完整可运行示例亦见仓库 examples/test_aclnn_right_shift.cpp具体编译与执行过程请参考 docs/zh/context/compile_and_run_sample.md。#include cstdint #include cstdio #include vector #include acl/acl.h #include aclnnop/aclnn_right_shift.h #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) int main() { constexpr int32_t deviceId 0; aclrtStream stream nullptr; auto ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtSetDevice(deviceId); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtCreateStream(stream); CHECK_RET(ret ACL_SUCCESS, return ret); const std::vectorint64_t shape {2, 4}; const std::vectorint32_t inputHost {-16, -8, -1, 0, 1, 8, 16, 32}; const std::vectorint32_t shiftHost {0, 1, 2, 3, -1, 32, 4, 5}; std::vectorint32_t outHost(inputHost.size(), 0); const size_t bytes inputHost.size() * sizeof(int32_t); void *inputDevice nullptr; void *shiftDevice nullptr; void *outDevice nullptr; ret aclrtMalloc(inputDevice, bytes, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtMalloc(shiftDevice, bytes, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtMalloc(outDevice, bytes, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtMemcpy(inputDevice, bytes, inputHost.data(), bytes, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtMemcpy(shiftDevice, bytes, shiftHost.data(), bytes, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret ACL_SUCCESS, return ret); std::vectorint64_t strides {4, 1}; aclTensor *input aclCreateTensor( shape.data(), shape.size(), ACL_INT32, strides.data(), 0, ACL_FORMAT_ND, shape.data(), shape.size(), inputDevice); aclTensor *shiftBits aclCreateTensor( shape.data(), shape.size(), ACL_INT32, strides.data(), 0, ACL_FORMAT_ND, shape.data(), shape.size(), shiftDevice); aclTensor *out aclCreateTensor( shape.data(), shape.size(), ACL_INT32, strides.data(), 0, ACL_FORMAT_ND, shape.data(), shape.size(), outDevice); uint64_t workspaceSize 0; aclOpExecutor *executor nullptr; ret aclnnRightShiftGetWorkspaceSize(input, shiftBits, out, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, return ret); void *workspace nullptr; if (workspaceSize 0) { ret aclrtMalloc(workspace, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret ACL_SUCCESS, return ret); } ret aclnnRightShift(workspace, workspaceSize, executor, stream); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtSynchronizeStream(stream); CHECK_RET(ret ACL_SUCCESS, return ret); ret aclrtMemcpy(outHost.data(), bytes, outDevice, bytes, ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET(ret ACL_SUCCESS, return ret); for (auto value : outHost) { std::printf(%d , value); } std::printf(\n); aclDestroyTensor(input); aclDestroyTensor(shiftBits); aclDestroyTensor(out); aclrtFree(inputDevice); aclrtFree(shiftDevice); aclrtFree(outDevice); if (workspace ! nullptr) { aclrtFree(workspace); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return ACL_SUCCESS; }示例数据刻意覆盖了典型边界场景便于对照验证文档语义移位位数为-1非法负数与32等于 INT32 位宽非法根据非法移位位数规则input为负数如-1时应输出-1非负数时应输出0其余合法移位按算术右移逐位计算。源码级原理aclnn 接口的内部计算流程第一段接口在完成参数校验后通过 Level0 算子接口l0op编排了一条完整的计算流水线见 op_api/aclnn_right_shift.cpp类型推导promoteop::PromoteType(input-GetDataType(), shiftBits-GetDataType())推导出统一的中间计算类型若推导失败结果为DT_UNDEFINED或推导结果无法 cast 成out的类型则返回ACLNN_ERR_PARAM_INVALID对应CheckPromoteType。连续化l0op::Contiguous将可能的非连续输入转换为连续 Tensor。类型统一l0op::Cast将input与shiftBits均转换为 promote 类型后送入计算。核心计算l0op::RightShift(inputCasted, shiftBitsdCasted, executor)生成右移计算节点。结果回写l0op::Cast将计算结果转换回out的数据类型再通过l0op::ViewCopy将结果拷贝到out上——这一步同时兼容了out为非连续 Tensor 的情形。workspace 获取uniqueExecutor-GetWorkspaceSize()汇总整条流水线所需 workspace通过workspaceSize返回执行器通过executor出参返回。第二段接口aclnnRightShift则直接调用CommonOpExecutorRun把第一段组装好的执行器提交到指定 Stream 上执行。源码级原理Kernel 的计算与非法移位处理Kernel 入口定义在 op_kernel/right_shift.cpp根据 tiling key 解出BROADCAST_MODE与DTYPE_MODE两个模板参数后分发到对应的RightShiftKernelImpl再按数据类型实例化为 8 种 Kernelint8_t 至 uint64_t见 op_kernel/right_shift.h。单个元素的计算逻辑ComputeOne严格遵循文档语义先经IsInvalidShift判断移位位数是否合法超过位宽减一或有符号类型下为负数非法时调用InvalidShiftValue有符号类型下xValue 0返回-1否则返回0无符号类型恒返回0合法时对有符号类型先扩展为int64_t再右移算术右移对无符号类型先提升为uint64_t再右移逻辑右移最后窄化回原类型。值得关注的是Kernel 对 16/32 位向量类型int16_t、uint16_t、int32_t、uint32_t对应IsVectorShiftType使用了桶式移位bucket shift优化以 256 位 repeat 为单位对每个可能的移位值shift0 到 bitWidth-1逐档执行CompareScalar生成掩码再用带掩码的ShiftRight向量指令批量移位移位位数非法元素的输出则通过InitInvalidShiftResult先行填好有符号类型用Maxs/Mins组合实现负数填 -1、非负填 0无符号类型直接填 0。对于 int8/uint8/int64/uint64 等无法走向量指令的类型则退化为逐元素标量计算。shiftBits为标量RIGHT_SHIFT_MODE_Y_SCALAR时还会走ProcessScalarY的标量移位快速路径。源码级原理Tiling 切分与 broadcast 模式Tiling 逻辑位于 op_host/right_shift_tiling.cpp负责把 shape、stride、核数与缓冲区长度等信息写入 op_kernel/right_shift_tiling_data.h 定义的RightShiftTilingData主要工作包括broadcast 信息构建对齐x/y的存储 shape 到相同 rank校验每个维度满足 broadcast 规则xDim yDim或其中一方为 1计算输出 shape 与总元素数并校验z的 shape 与 broadcast 结果一致。broadcast 压缩将连续相同 broadcast 状态的维度合并CompressBroadcastInfo并计算出xStride/yStride维度为 1 且需要扩展的维度 stride 置 0供 Kernel 侧做地址映射。broadcast 模式分类根据形状关系划分为 5 种模式见 op_kernel/right_shift_tiling_data.hRIGHT_SHIFT_MODE_CONTIGUOUS0x、y与输出完全同 shape直接连续搬运RIGHT_SHIFT_MODE_X_SCALAR1x为标量元素数为 1RIGHT_SHIFT_MODE_Y_SCALAR2y为标量走标量移位快速路径RIGHT_SHIFT_MODE_TAIL_CONTIGUOUS3最内层维度连续且可整块搬运RIGHT_SHIFT_MODE_GENERAL4一般 broadcast 场景按段拷贝与填充。资源与核数规划通过PlatformAscendC查询 UB 大小与可用 AIV 核数按每元素占用 5 倍类型宽度 4 字节比较缓冲估算 tile 长度RIGHT_SHIFT_TMP_BUFFER_FACTOR 5并按总元素数决定单核/多核划分总长不超过 2048 时使用单核见MULTI_CORE_SIZE_LIMIT最终以formerCoreNum tailCoreNum设置 block 维度并把mode * 8 dtypeIndex编码进 tiling key。测试用例参考仓库为该算子提供了三层 UTop_api 层tests/ut/op_api/test_aclnn_right_shift.cpp 覆盖同 shape、y标量、x标量、多维 broadcast、不同数据类型组合以及非法入参错误返回码等场景例如case_01_int32_same_shape{2,3}同 shape、case_02_uint8_y_scalar、case_03_int64_x_scalar、case_04_int16_broadcast{2,1,4}与{1,3,4}广播为{2,3,4}。op_host 层tests/ut/op_host/test_right_shift_tiling.cpp 针对 tiling 函数验证 broadcast 模式判定、stride 计算与核数划分。op_kernel 层tests/ut/op_kernel/test_right_shift.cpp 结合 tests/ut/op_kernel/CMakeLists.txt 构建运行直接验证 Kernel 的计算正确性。开发者如需快速体验可参考 examples/test_aclnn_right_shift.cpp 作为最小可运行样例并配合仓库根目录 README.md 与 CONTRIBUTING.md 了解整体构建与贡献流程。赞分享算子库人工智能CANN【免费下载链接】ops-math本项目是CANN提供的数学类基础计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-math点击查看免费下载相关推荐JSON for Modern C 深度解析byte_container_with_subtype 构造函数与二进制子类型机制JSON for Modern C 深度解析byte_container_with_subtype 构造函数与二进制子类型机制 本文围绕 nlohmann算子库人工智能CANNCANN ops-math Sqrt 算子 aclnn 接口使用指南两段式调用、参数约束与源码实现解析CANN ops math Sqrt 算子 aclnn 接口使用指南两段式调用、参数约束与源码实现解析 导读 本文以 CANN ops math 仓库中 Sq算子库人工智能CANNCANN ops-math 算子解读aclnnCummax 两段式接口原理、参数约束与调用示例CANN ops math 算子解读aclnnCummax 两段式接口原理、参数约束与调用示例 aclnnCummax 是 CANN ops math 数学算算子库人工智能CANN创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考