pyasc 的 asc.language.basic.max 深度解析:LocalTensor 逐元素最大值与高维切分计算
【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc
asc.language.basic.max是 pyasc 中面向昇腾 AI 处理器 Vector Core 的逐元素二元算子接口,用于在LocalTensor之间按元素求最大值,语义与 Ascend C 的Max函数一一对应。本文基于当前仓库中的 API 文档、源码实现与单元测试,完整讲解该接口的三种重载形式、参数含义与取值约定、数据类型约束,以及从 Python 调用到 Ascend C 代码生成的底层链路,帮助你在用 Python 编写向量算子时正确、高效地完成元素级比较计算与高维张量切分迭代。
接口定位:Python 原语对应 Ascend C Max
在 Ascend C 体系里,Max是一组作用于 UB(统一缓冲)上LocalTensor的二元矢量指令。pyasc 将其包装为符合 Python 原生语法的asc.max,通过 JIT 方式在编译期把 Python 调用翻译成对应的 IR 操作,并最终生成与手写 Ascend C 等价的ascendc::Max(...)调用。
接口文档中给出了它对应的 Ascend C 函数原型,共三种形态:
template <typename T> __aicore__ inline void Max(const LocalTensor<T>& dst, const LocalTensor<T>& src0, const LocalTensor<T>& src1, const int32_t& count);template <typename T, bool isSetMask = true> __aicore__ inline void Max(const LocalTensor<T>& dst, const LocalTensor<T>& src0, const LocalTensor<T>& src1, uint64_t mask[], const uint8_t repeatTimes, const BinaryRepeatParams& repeatParams);template <typename T, bool isSetMask = true> __aicore__ inline void Max(const LocalTensor<T>& dst, const LocalTensor<T>& src0, const LocalTensor<T>& src1, uint64_t mask, const uint8_t repeatTimes, const BinaryRepeatParams& repeatParams);这三种 C++ 原型分别对应 pyasc 中三个按参数特征区分的中载签名(见 API 文档):
asc.language.basic.max(dst, src0, src1, count, is_set_mask=True)—— 以元素个数count指定运算量,对应"tensor 前 n 个数据计算"场景;asc.language.basic.max(dst, src0, src1, mask: int, repeat_times, repeat_params, is_set_mask=True)—— mask 为连续模式,一次迭代连续处理 mask 个元素,适合高维切分迭代;asc.language.basic.max(dst, src0, src1, mask: List[int], repeat_times, repeat_params, is_set_mask=True)—— mask 为逐 bit 模式,掩码数组的每个 bit 决定是否处理对应元素。
三者返回值均为None:计算结果直接写回dst,接口本身不产生新对象。
源码中的重载分发
从源码看,max定义在 vec_binary.py:
@overload def max(dst: LocalTensor, src0: LocalTensor, src1: LocalTensor, count: int, is_set_mask: bool = True) -> None: ... @overload def max(dst: LocalTensor, src0: LocalTensor, src1: LocalTensor, mask: int, repeat_times: int, repeat_params: BinaryRepeatParams, is_set_mask: bool = True) -> None: ... @overload def max(dst: LocalTensor, src0: LocalTensor, src1: LocalTensor, mask: List[int], repeat_times: int, repeat_params: BinaryRepeatParams, is_set_mask: bool = True) -> None: ... @require_jit @set_binary_docstring(cpp_name="Max", append_text="按元素求最大值。") def max(dst: LocalTensor, src0: LocalTensor, src1: LocalTensor, *args, **kwargs) -> None: builder = global_builder.get_ir_builder() op_impl("max", dst, src0, src1, args, kwargs, builder.create_asc_MaxL0Op, builder.create_asc_MaxL1Op, builder.create_asc_MaxL2Op)几个值得注意的实现细节:
@require_jit表明该接口只能在 JIT 编译上下文中调用,即函数体处于 pyasc 的算子编译流程内;- 实际执行走通用的二元算子分发器 op_impl,按关键字参数类型匹配三条路径:
mask为整数(RuntimeInt)→ 构造asc.MaxL0Op(连续模式),mask 物化为 64 位整型 IR 值;mask为列表 → 逐元素物化为 uint64 后构造asc.MaxL1Op(逐 bit 模式);- 只传
count→ 构造asc.MaxL2Op,count 物化为 int32;
is_set_mask默认True,会被一并传入 IR 操作,决定 C++ 模板参数isSetMask的取值。
也就是说,用户写 Python 重载 2/3 时,pyasc 在编译期会生成MaxL0Op/MaxL1Op,is_set_mask参数则直接映射到 C++ 侧的模板布尔参数上。
参数说明
| 参数 | 类型 | 说明 |
|---|---|---|
dst | LocalTensor | 目的操作数,计算结果写回此张量。支持的 TPosition 为VECIN/VECCALC/VECOUT。 |
src0,src1 | LocalTensor | 源操作数,两两配对比较。支持的 TPosition 为VECIN/VECCALC/VECOUT。 |
count | int | 参与计算的元素个数,仅第一种重载使用。 |
mask | int或List[int] | 控制每次迭代内参与计算的元素:整数为连续模式(一次迭代连续处理 mask 个元素),列表为逐 bit 模式(掩码的每个 bit 对应一个元素)。 |
repeat_times | int | 重复迭代次数。总处理量 = mask 控制量 × repeat_times。 |
repeat_params | BinaryRepeatParams | 控制三个操作数地址步长的参数,决定迭代内(blk_stride)与迭代间(rep_stride)的地址推进方式。 |
is_set_mask | bool,默认True | 是否在接口内部设置 mask。False表示 mask 在接口外部设置(如先调用set_vector_mask),接口内部不再重复设置。 |
其中LocalTensor是 UB 上的张量抽象,详见 core 模块文档。
BinaryRepeatParams:步长参数与默认值
高维切分场景的核心是BinaryRepeatParams,它由六个步长组成。从 types.py 中的定义看,默认值为:
BinaryRepeatParams( dst_blk_stride=1, # 单次迭代内 dst 各 data block 的地址步长 src0_blk_stride=1, # 单次迭代内 src0 各 data block 的地址步长 src1_blk_stride=1, # 单次迭代内 src1 各 data block 的地址步长 dst_rep_stride=8, # 相邻迭代之间 dst 的地址步长 src0_rep_stride=8, # 相邻迭代之间 src0 的地址步长 src1_rep_stride=8, # 相邻迭代之间 src1 的地址步长 )*_blk_stride控制"一次迭代内部"各数据块之间的间距,取 1 表示迭代内数据连续读写;*_rep_stride控制"相邻迭代之间"地址的推进量,默认 8 表示迭代间同样保持连续衔接。
两者组合起来可以描述任意"块内连续、块间按固定步长跳跃"的二维(乃至更高维展开)布局,这正是处理高维张量切分后非连续内存排布的关键。
数据类型约束
pyasc 在分发器入口处对操作数类型做静态校验。utils.py 中check_type为"max"登记的合法类型是:
valids = {"src": [KT.float16, KT.float32, KT.int16, KT.int32], "dst": [KT.float16, KT.float32, KT.int16, KT.int32]}且校验逻辑强制src0、src1类型一致,并且dst与两个源类型必须完全相同(max不在允许类型转换的接口集合中)。因此调用时若传入如float32的 dst 配float16的 src,会直接抛出TypeError,在编译期就暴露问题,而非运行时报错。max与其同族的add、min、mul、sub共享同一套类型集合。
调用示例
以下示例完整继承自 API 文档,并标注了参数取值含义。
场景一:高维切分计算,mask 连续模式
mask = 128 # repeat_times = 4,一次迭代计算128个数,共计算512个数 # dst_blk_stride, src0_blk_stride, src1_blk_stride = 1,单次迭代内数据连续读取和写入 # dst_rep_stride, src0_rep_stride, src1_rep_stride = 8,相邻迭代间数据连续读取和写入 params = asc.BinaryRepeatParams(1, 1, 1, 8, 8, 8) asc.max(dst, src0, src1, mask=mask, repeat_times=4, repeat_params=params)这里mask=128表示单次迭代内连续比较 128 个元素,重复 4 次共处理 512 个元素;步长配置(1,1,1,8,8,8)意味着迭代内块间连续、迭代间也连续衔接,等价于对一段 512 元素区域做分块扫描。
场景二:高维切分计算,mask 逐 bit 模式
mask = [uint64_max, uint64_max] # uint64_max = 2**64 - 1 # repeat_times = 4,一次迭代计算128个数,共计算512个数 params = asc.BinaryRepeatParams(1, 1, 1, 8, 8, 8) asc.max(dst, src0, src1, mask=mask, repeat_times=4, repeat_params=params)逐 bit 模式下掩码是 uint64 数组,每个 bit 对应一个元素位(16 bit 数据下两个 uint64 恰好覆盖 128 个元素)。两个0xFFFFFFFFFFFFFFFF表示本迭代内所有位都参与计算;若只想处理交错位置,可以把某些 bit 清零,实现比"连续段"更细粒度的元素选择。
场景三:tensor 前 n 个数据计算
asc.max(dst, src0, src1, count=512)最简单的形态:从张量起始地址开始比较前 512 个元素。使用整个 tensor 参与计算(即不显式切分)时,运算量就是目的LocalTensor的总长度。
仓库单元测试 test_vector_binary.py 中对三种形态的用法与文档一致,可作为可运行的参照:
x_local = asc.LocalTensor(dtype=asc.float16, pos=asc.TPosition.VECIN, addr=0, tile_size=512) y_local = asc.LocalTensor(dtype=asc.float16, pos=asc.TPosition.VECIN, addr=0, tile_size=512) z_local = asc.LocalTensor(dtype=asc.float16, pos=asc.TPosition.VECOUT, addr=0, tile_size=512) asc.max(z_local, x_local, y_local, count=512) params = asc.BinaryRepeatParams(1, 1, 1, 8, 8, 8) asc.max(z_local, x_local, y_local, mask=512, repeat_times=1, repeat_params=params) uint64_max = 2**64 - 1 asc.max(z_local, x_local, y_local, mask=[uint64_max, uint64_max], repeat_times=1, repeat_params=params)可以看到在 fp16、512 元素的场景下,连续模式用mask=512, repeat_times=1一次迭代完成,逐 bit 模式用两个满掩码 uint64 覆盖同样范围——三种重载处理的数据量应当一致,只是掩码表达粒度不同。
底层实现链路:从 Python 调用到 C++ 代码生成
pyasc 的整体定位是"Python 写算子、编译出昇腾可执行体",max的完整链路如下(整体架构参见 pyasc_arch):
- Python 前端:vec_binary.py 中
max在 JIT 上下文中拿到 IR builder,调用op_impl; - 类型校验与分发:utils.py 先执行
check_type,再按参数形态注册到OverloadDispatcher,选择create_asc_MaxL0Op/create_asc_MaxL1Op/create_asc_MaxL2Op之一构造 AscIR 操作,并把mask、repeat_times物化为具体位宽的 IR 标量(mask 为 int64/uint64 列表,repeat_times 为 int8,count 为 int32); - AscIR 方言:三种操作在 AscIR 方言中登记(MaxL0Op/MaxL1Op/MaxL2Op 在 Translation.cpp 的二元算子注册列表中可见),携带
dst/src0/src1/mask/repeatTimes/repeatParams/isSetMask等字段; - 代码发射:VecBinary.h 中的打印模板负责把 IR 还原为 C++ 调用文本——L2 形态发射
ascendc::Max(dst, src0, src1, count),L0/L1 形态先打印isSetMask模板参数再输出(dst, src0, src1, mask, repeatTimes, repeatParams)实参列表。生成的 C++ 与手写 Ascend C 的Max调用形式一致,因此后续交由标准昇腾工具链编译,行为与 Ascend C 原语完全对齐。
从源码结构看,max、min、add、mul、sub等二元算子共用同一套op_impl分发器与打印模板,max的差异仅在于传给 builder 的工厂函数(create_asc_Max*Op)与校验表中的类型集合。这一设计意味着理解max即掌握了 pyasc 全部逐元素二元算子的通用调用模式。
约束说明
- 地址对齐:操作数的地址对齐要求遵循《Ascend C算子开发接口》中"通用说明和约束-通用地址对齐约束"一节的说明,需按所支持数据类型满足相应字节对齐。
- 地址重叠:操作数之间的地址重叠约束同样参见《Ascend C算子开发接口》中"通用说明和约束-通用地址重叠约束"。若
dst与src0/src1存在重叠,是否安全取决于该文档中给出的规则,不应假设原地操作一定正确。 - 运算量约定:使用整个 tensor 参与计算(
count形态)时,运算量为目的LocalTensor的总长度;mask 形态下运算量由mask与repeat_times共同决定,应确保不超过张量实际长度。 - is_set_mask 与外部 mask:当
is_set_mask=False时,mask 需在接口外部先行设置(例如通过矢量掩码设置接口),接口内部不再覆盖该值;此时 mask 参数应作占位理解,避免内外两层设置互相冲突。 - 适用前提:以上均针对当前仓库版本的 pyasc,接口需在 JIT 编译上下文(
@require_jit约束)中调用;具体硬件能力以实际昇腾芯片的 Ascend C 接口文档为准。
小结
asc.language.basic.max是 pyasc 对 Ascend CMax原语的完整 Python 化映射:count形态覆盖前 n 个元素的直白计算,mask连续/逐 bit 两种形态配合BinaryRepeatParams的六步长参数则支撑高维张量的分块切分迭代。仓库中 API 文档、实现代码、类型校验、IR 到 C++ 的打印模板 与单元测试 相互印证,构成从使用到实现的完整证据链;同类接口(如 min、add)可参考本文的解析路径继续查阅。
【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考