ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

CANN pyasc 按元素开方算子接口 `asc.language.basic.sqrt` 完全指南:重载形式、mask 迭代计算与端到端用例解析

CANN pyasc 按元素开方算子接口 `asc.language.basic.sqrt` 完全指南:重载形式、mask 迭代计算与端到端用例解析 CANN pyasc 按元素开方算子接口asc.language.basic.sqrt完全指南重载形式、mask 迭代计算与端到端用例解析【免费下载链接】pyasc本项目为Python用户提供算子编程接口支持在昇腾AI处理器上加速计算接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc本文围绕 CANN/pyasc 仓库中 docs/python-api/language/generated/asc.language.basic.sqrt.md 所定义的asc.language.basic.sqrt接口展开系统讲解其在昇腾 AI 处理器上对LocalTensor按元素执行开方Sqrt计算的三种重载形式、mask 连续/逐 bit 两种切分计算模式、UnaryRepeatParams步长参数语义并结合仓库 Python 源码、后端翻译注册与端到端泛化测试说明该接口从 Python 调用到 Ascend C 原生Sqrt函数的完整落盘路径。读完本文你将能够独立使用asc.sqrt编写可多核并行、带 tiling 切片的向量开方算子并理解count与mask/repeat_times两套编程范式的取舍。接口定位与功能概述asc.language.basic.sqrt是 CANN pyasc 提供的按元素做开方向量一元运算接口对应原生 Ascend C 编程接口中的Sqrt。它接收两个LocalTensor目的操作数dst与源操作数src将src中每个元素的开方结果写入dst对应位置即逐元素执行dst[i] sqrt(src[i])。在 python/asc/language/basic/init.py 中sqrt与abs、exp、ln、reciprocal、relu、rsqrt等一同作为向量一元运算Unary对外暴露其 Python 实现位于 python/asc/language/basic/vec_unary.py。接口的函数签名、文档字符串由 python/asc/language/basic/utils.py 中的set_unary_docstring(cpp_nameSqrt, append_text按元素做开方。)装饰器统一生成这也解释了为什么该接口的在线文档与源码 docstring 结构完全一致。三种重载形式与参数说明asc.language.basic.sqrt支持三种重载overload分别面向前 N 个元素整段计算与高维切分迭代计算两类场景# 形式一按元素个数整体计算 asc.language.basic.sqrt(dst: LocalTensor, src: LocalTensor, count: int) - None # 形式二mask 为单个整数的连续模式 asc.language.basic.sqrt(dst: LocalTensor, src: LocalTensor, mask: int, repeat_times: int, repeat_params: UnaryRepeatParams, is_set_mask: bool True) - None # 形式三mask 为整数列表的逐 bit 模式 asc.language.basic.sqrt(dst: LocalTensor, src: LocalTensor, mask: List[int], repeat_times: int, repeat_params: UnaryRepeatParams, is_set_mask: bool True) - None对应地三种重载映射到三种 Ascend C 函数原型见原文档// 与形式一对应calCount 参与计算的元素个数 template typename T __aicore__ inline void Sqrt(const LocalTensorT dstLocal, const LocalTensorT srcLocal, const int32_t calCount) // 与形式二对应mask 为 uint64_t 标量 template typename T, bool isSetMask true __aicore__ inline void Sqrt(const LocalTensorT dstLocal, const LocalTensorT srcLocal, uint64_t mask, const uint8_t repeatTimes, const UnaryRepeatParams repeatParams) // 与形式三对应mask 为 uint64_t 数组 template typename T, bool isSetMask true __aicore__ inline void Sqrt(const LocalTensorT dstLocal, const LocalTensorT srcLocal, uint64_t mask[], const uint8_t repeatTimes, const UnaryRepeatParams repeatParams)各参数语义如下dst目的操作数类型为LocalTensor支持的 TPosition 为VECIN/VECCALC/VECOUT。LocalTensor的接口定义可参考 LocalTensor 文档源码实现见 python/asc/language/core/tensor.py。src源操作数类型同样为LocalTensorTPosition 约束与dst一致。count参与计算的元素个数int仅在形式一中使用接口内部根据元素个数完成整段开方。mask控制每次迭代repeat内参与计算的元素。形式二中为单个int连续 bit 模式形式三中为List[int]逐 bit 模式。repeat_times重复迭代次数即整段数据按 mask 粒度切分后需要执行的迭代轮数。repeat_paramsUnaryRepeatParams类型控制操作数在单次迭代内blk 级与相邻迭代间rep 级的地址步长。is_set_mask是否在接口内部设置 mask默认True。传False时由调用方预先通过asc.set_mask_count/asc.set_vector_mask等接口设置好 mask接口不再重复下发。从 Python 重载分发到 IR 建图L0/L1/L2 三级实现sqrt的 Python 侧实现是理解三种重载底层行为的关键。见 python/asc/language/basic/vec_unary.py 中所有一元运算共用的op_impl分发逻辑overload def sqrt(dst: LocalTensor, src: LocalTensor, count: int) - None: ... overload def sqrt(dst: LocalTensor, src: LocalTensor, mask: int, repeat_times: int, repeat_params: UnaryRepeatParams, is_set_mask: bool True) - None: ... overload def sqrt(dst: LocalTensor, src: LocalTensor, mask: List[int], repeat_times: int, repeat_params: UnaryRepeatParams, is_set_mask: bool True) - None: ... require_jit set_unary_docstring(cpp_nameSqrt, append_text按元素做开方。) def sqrt(dst: LocalTensor, src: LocalTensor, *args, **kwargs) - None: builder global_builder.get_ir_builder() op_impl(sqrt, dst, src, args, kwargs, builder.create_asc_SqrtL0Op, builder.create_asc_SqrtL1Op, builder.create_asc_SqrtL2Op)op_impl内部通过OverloadDispatcher按参数形态完成重载判定并分别建图mask 为int形式二→ 调用builder.create_asc_SqrtL0Op构建L0级 IR。其中 mask 会被物化为uint64标量repeat_times物化为int8is_set_mask作为模板布尔参数传入。mask 为list形式三→ 调用builder.create_asc_SqrtL1Op构建L1级 IR。列表中的每个元素先经materialize_ir_value物化为uint64构成uint64_t mask[]数组对应逐 bit 模式。仅传count形式一→ 走dispatcher.register_auto兜底分支调用builder.create_asc_SqrtL2Op构建L2级 IRcount物化为int32。L0/L1/L2 三种 IR 算子ascendc::SqrtL0Op/SqrtL1Op/SqrtL2Op在昇腾后端统一注册为向量一元运算见 lib/Target/AscendC/Translation.cpp 中// Vector unary operations一节的注册列表同一文件第 79-88 行还将SqrtRegOp注册为寄存器级一元运算支撑标量/寄存器场景。此外math::SqrtOp也被纳入外部 Math 方言翻译Translation.cpp说明 sqrt 语义在 pyasc 的 IR 体系中具备完整的从 MLIR 方言到 Ascend C 文本的翻译链。可见 pyasc 的设计是一层 Python 语法、三种底层建图路径对用户而言始终是asc.sqrt(dst, src, ...)一个入口而 mask 的形态标量 vs 列表与是否使用 count决定了最终落到哪条向量指令路径。mask 机制与UnaryRepeatParams步长语义mask与repeat_params共同决定一次迭代算多少元素、数据如何寻址是向量指令编程的核心。昇腾向量单元的典型做法是把一段连续数据按 repeat迭代切分每次迭代内再按 mask 决定参与运算的元素子集。UnaryRepeatParams的构造定义见 python/asc/language/core/types.pyrequire_jit def __init__(self, dst_blk_stride: RuntimeInt 1, src_blk_stride: RuntimeInt 1, dst_rep_stride: RuntimeInt 8, src_rep_stride: RuntimeInt 8, ...): # 通过 builder.create_asc_ConstructOp 构造 UnaryRepeatParams 类型的 IR 值四个字段的工程含义如下默认值已在源码中确认字段默认值含义dst_blk_stride1目的操作数在单次迭代内相邻 block 间的步长单位blocksrc_blk_stride1源操作数在单次迭代内相邻 block 间的步长dst_rep_stride8目的操作数在相邻迭代repeat间的步长单位blocksrc_rep_stride8源操作数在相邻迭代repeat间的步长其中blk_stride 1表示单次迭代内数据连续读写rep_stride 8表示相邻迭代间各跳过 8 个 block 再继续读写配合连续的 mask 可达到整段连续搬运的效果。需要说明的是UnaryRepeatParams被物化为asc_UnaryRepeatParamsTypeIR 类型其中dst_blk_stride/src_blk_stride为uint16dst_rep_stride/src_rep_stride为uint8传入超出位宽的值会在建图阶段被约束因此实际使用时建议遵循默认量级。调用示例完整继承原文档原文档给出三类可直接落地的示例以下完整保留并补充注释。1. mask 连续模式单 int maskmask 256 // asc.half.sizeof()以asc.half为例half占 2 字节256 // 2 128即每次迭代通过连续 bit mask 选中 128 个元素repeat_times 4共计算4 × 128 512个数。dst_blk_stride/src_blk_stride 1保证单次迭代内连续读写dst_rep_stride/src_rep_stride 8保证相邻迭代间连续读写每次迭代消费 8 个 block迭代间无缝衔接。mask 256 // asc.half.sizeof() # repeat_times 4一次迭代计算128个数共计算512个数 # dst_blk_stride, src_blk_stride 1单次迭代内数据连续读取和写入 # dst_rep_stride, src_rep_stride 8相邻迭代间数据连续读取和写入 params asc.UnaryRepeatParams(1, 1, 8, 8) asc.sqrt(dst, src, maskmask, repeat_times4, repeat_paramsparams)2. mask 逐 bit 模式List mask通过两个uint64_max组成 128 bit 全 1 掩码同样实现每次迭代计算 128 个数、4 次迭代共 512 个数的效果区别在于 mask 以逐 bit 列表形式显式给出适用于需要精确控制每个元素参与与否的场景例如非对齐或需跳过特定元素的高维切分。mask [uint64_max, uint64_max] # repeat_times 4一次迭代计算128个数共计算512个数 # dst_blk_stride, src_blk_stride 1单次迭代内数据连续读取和写入 # dst_rep_stride, src_rep_stride 8相邻迭代间数据连续读取和写入 params asc.UnaryRepeatParams(1, 1, 8, 8) asc.sqrt(dst, src, maskmask, repeat_times4, repeat_paramsparams)3. 前 N 个元素整段计算count 模式当数据在本地内存中连续、无需精细切分时直接传入元素个数即可由 L2 级指令内部完成整段处理asc.sqrt(dst, src, count512)约束说明使用asc.language.basic.sqrt时需遵守与 Ascend C 一元向量算子一致的两类通用约束操作数地址对齐约束dst与src的起始地址需满足向量单元的数据对齐要求一般按 block/32B 粒度对齐否则可能导致指令非法或结果错误。详见《Ascend C 算子开发接口》中的通用说明和约束-通用地址对齐约束。操作数地址重叠约束dst与src指向的本地内存区间原则上不应重叠或重叠方式必须符合接口约定否则计算结果不可预期。详见《Ascend C 算子开发接口》中的通用说明和约束-通用地址重叠约束。此外从实现看vec_unary.pyrepeat_times会被物化为int8、count物化为int32、mask 元素物化为uint64因此这三个入参的实际取值范围受对应位宽约束。端到端验证多核 tiling 的完整算子形态仅了解单接口调用还不够仓库中的泛化测试 python/test/generalization/basic/test_vsqrt.py 给出了asc.sqrt在真实算子kernel中的完整用法是本文接口在生产形态下的最佳范本。该测试展示了标准的数据流式算子骨架多核切分use_core_num 16每个核通过asc.get_block_idx() * block_length计算自己的数据段偏移流水队列asc.TQue(asc.TPosition.VECIN, buffer_num)与asc.TQue(asc.TPosition.VECOUT, buffer_num)分别承载输入与输出TPipe.init_buffer分配双缓冲/多缓冲三段式循环copy_inasc.data_copy从 Global 搬入→compute核心的asc.sqrt(z_local, x_local, counttile_length)→copy_out结果写回 Globaltiling 参数tile_length 512tile_num ceil(block_length / tile_length)对应 count 模式的整段计算正确性校验与torch.sqrt逐元素allclose对比float16 容忍rtol1e-3float32 使用默认精度。其中 compute 阶段的核心调用正是本文主角test_vsqrt.pyasc.jit def compute(z_gm, in_queue_x, out_queue_z, tile_length): x_local in_queue_x.deque(z_gm.dtype) z_local out_queue_z.alloc_tensor(z_gm.dtype) asc.sqrt(z_local, x_local, counttile_length) out_queue_z.enque(z_local) in_queue_x.free_tensor(x_local)测试覆盖的用例规模包括float16 (2048,)、float32 (5000,)、float32 (9999,)、float32 (8192,)等非对齐/对齐混合形状并通过config.Backend.NPU在真实 NPU 平台运行test_vsqrt.py。此外单元测试 python/test/unit/language/basic/test_vector_unary.py 覆盖了接口的建图与参数合法性校验。这意味着如果你要在一个自定义算子中使用开方完全可以直接参照vsqrt_kernel的骨架把asc.sqrt嵌入你自己的 tiling 循环。小结两套编程范式的选择综合原文档与源码使用asc.sqrt时的决策路径可以归纳为数据段连续、只需对整段开方如每个 tile 恰好 512 个连续元素→ 使用count 模式形式一代码最简洁也是test_vsqrt采用的形态需要按 128/256 元素粒度迭代计算、数据在 block 间连续但需显式控制迭代步长 → 使用mask 连续模式形式二mask: intUnaryRepeatParams需要对每个 bit 精确控制参与元素非对齐、跳元素、不规则高维切分→ 使用mask 逐 bit 模式形式三mask: List[int]若已自行通过set_mask_count/set_vector_mask设置过 mask可令is_set_maskFalse避免接口重复设置。三者共享同一 Python 入口与dst/src的 TPosition 约束最终都经由 L0/L1/L2 三种 IR 之一翻译为昇腾向量单元上的原生Sqrt指令是 CANN pyascPython 原生语法 Ascend C 一一对应设计理念在向量一元运算上的典型体现。【免费下载链接】pyasc本项目为Python用户提供算子编程接口支持在昇腾AI处理器上加速计算接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进