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

资讯详情

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

CANN pyasc 融合乘加算子 fused_mul_add 完全指南:从接口语义到向量流水线底层实现

CANN pyasc 融合乘加算子 fused_mul_add 完全指南:从接口语义到向量流水线底层实现 CANN pyasc 融合乘加算子 fused_mul_add 完全指南从接口语义到向量流水线底层实现【免费下载链接】pyasc本项目为Python用户提供算子编程接口支持在昇腾AI处理器上加速计算接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc本文以 CANN pyasc 仓库中的asc.language.basic.fused_mul_add官方 API 文档docs/python-api/language/generated/asc.language.basic.fused_mul_add.md为骨架结合 vec_binary.py、utils.py、types.py 等源码与单元测试系统讲解这一融合乘加向量指令的三种调用形态count 连续模式 / mask 连续模式 / mask 逐 bit 模式、BinaryRepeatParams步长控制参数、数据类型约束与地址对齐要求。读完本文你将能够在昇腾 AI 处理器上通过 Python 原生语法正确编写dst src0 * dst src1的融合计算算子并理解它背后从 pyasc 前端到 Ascend C 模板函数再到 MLIR 算子的完整调用链路。接口签名与计算语义fused_mul_add是 pyasc 提供的向量二元融合计算接口对应 Ascend C 中的FusedMulAdd模板函数。其核心语义是按元素将 src0 与 dst 相乘再加上 src1最终结果写回 dstdst[i] src0[i] * dst[i] src1[i]该接口位于asc.language.basic模块可通过asc.fused_mul_add(...)直接调用。pyasc 通过overload声明了三种重载形态在运行时由OverloadDispatcher依据实参形态自动分发分别对应底层 L0 / L1 / L2 三种指令重载形态触发条件底层指令适用场景count 模式传入countL2create_asc_FusedMulAddL2Op整段连续数据接口内部按目的 tensor 长度自动计算迭代与步长mask 整数模式传入mask: intL0create_asc_FusedMulAddL0Opmask 连续模式一次迭代内用单个 uint64 掩码控制参与计算的元素mask 列表模式传入mask: List[int]L1create_asc_FusedMulAddL1Opmask 逐 bit 模式用 uint64 列表逐 bit 控制每次迭代内元素三种形态的函数原型如下asc.fused_mul_add(dst, src0, src1, count, is_set_maskTrue) # count 模式 asc.fused_mul_add(dst, src0, src1, mask, repeat_times, repeat_params, is_set_maskTrue) # mask: int asc.fused_mul_add(dst, src0, src1, mask, repeat_times, repeat_params, is_set_maskTrue) # mask: List[int]对应的 Ascend C 函数原型见 vec_binary.py 中的 docstring 生成逻辑template typename T __aicore__ inline void FusedMulAdd(const LocalTensorT dst, const LocalTensorT src0, const LocalTensorT src1, const int32_t count); template typename T, bool isSetMask true __aicore__ inline void FusedMulAdd(const LocalTensorT dst, const LocalTensorT src0, const LocalTensorT src1, uint64_t mask, const uint8_t repeatTimes, const BinaryRepeatParams repeatParams); template typename T, bool isSetMask true __aicore__ inline void FusedMulAdd(const LocalTensorT dst, const LocalTensorT src0, const LocalTensorT src1, uint64_t mask[], const uint8_t repeatTimes, const BinaryRepeatParams repeatParams);参数说明dst目的操作数类型为LocalTensor支持 TPosition 为VECIN/VECCALC/VECOUT。计算结果为src0 * dst src1并写回 dst。src0, src1源操作数类型为LocalTensorTPosition 同样支持VECIN/VECCALC/VECOUT。count参与计算的元素个数仅在 count 模式下使用。此时运算量为目的LocalTensor的总长度接口按整个 tensor 参与计算。mask控制每次迭代repeat内参与计算的元素。整数形式为 uint64 位图列表形式为 uint64 数组每个元素对应一次迭代的逐 bit 掩码。repeat_times重复迭代次数即该向量指令被重复执行的轮数。repeat_params控制操作数地址步长的参数类型为BinaryRepeatParams用于刻画数据在 block 级与 repeat 级的排布。is_set_mask是否在接口内部设置 mask默认True。对应 Ascend C 模板参数isSetMask。BinaryRepeatParams 步长参数详解BinaryRepeatParams定义于 types.py其构造函数与默认值如下asc.BinaryRepeatParams(dst_blk_stride1, src0_blk_stride1, src1_blk_stride1, dst_rep_stride8, src0_rep_stride8, src1_rep_stride8)dst_blk_stride / src0_blk_stride / src1_blk_strideblock 级步长即单次迭代一个 repeat内部相邻 32B block 之间的地址间隔默认均为 1表示单次迭代内数据连续读取和写入。dst_rep_stride / src0_rep_stride / src1_rep_striderepeat 级步长即相邻两次迭代起始地址之间的间隔默认均为 8block 数表示相邻迭代间数据连续读取和写入。从源码实现看BinaryRepeatParams构造时会通过builder.create_asc_ConstructOp构建一个asc_BinaryRepeatParamsType的 IR 值六个参数均按 uint8 类型builder.get_ui8_type()编码进构造指令中最终在向量流水线上驱动 DMA 地址生成。步长单位是 32B block因此在连续数据的朴素场景下保持默认值即可只有在做高维 tensor 切分、stride 访问时才需要显式调整。数据类型与合法性约束fused_mul_add属于浮点融合运算其数据类型约束在 utils.py 的check_type中由valids_float定义valids_float {src: [KT.float16, KT.float32], dst: [KT.float16, KT.float32]}即dst / src0 / src1 均只支持float16与float32src0 与 src1 必须同类型且由于fused_mul_add不在check_dst_src豁免集合中dst 必须与 src0、src1 类型完全一致违反上述约束时check_type会抛出带期望类型提示的TypeError。对比同文件中的其他算子可以更清晰地把握其定位add、max、min、mul等支持float16/float32/int16/int32四类而fused_mul_add、fused_mul_add_relu、div、mul_add_dst等融合/除法算子仅支持浮点类型。通用地址约束地址对齐操作数地址对齐要求遵循《Ascend C 算子开发接口》中通用说明和约束-通用地址对齐约束。地址重叠操作数地址重叠约束参考通用说明和约束-通用地址重叠约束。需注意 dst 同时参与乘法与写回若 dst 与 src0/src1 地址重叠其结果语义按硬件流水线实际读写顺序确定应避免未定义的别名访问。运算量使用整个 tensor 参与计算count 模式接口符号重载时运算量为目的LocalTensor的总长度。调用示例三种使用形态以下示例完整继承自官方 API 文档并补充了参数语义注释可直接在asc.jit装饰的核函数中运行。1. tensor 高维切分计算样例 —— 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.fused_mul_add(dst, src0, src1, maskmask, repeat_times4, repeat_paramsparams)mask 为单个整数128表示每次迭代内只计算前 128 个元素uint64 掩码的 bit 置位连续迭代 4 次共计算 512 个数。2. tensor 高维切分计算样例 —— 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.fused_mul_add(dst, src0, src1, maskmask, repeat_times4, repeat_paramsparams)mask 为 uint64 列表时进入逐 bit 模式[uint64_max, uint64_max]表示每次迭代使用两个 64-bit 掩码共覆盖 128 个元素每个 bit 对应一个元素配合repeat_times4同样完成 512 个元素的融合乘加。该形态允许每次迭代内部只计算非连续的子集是高维 tensor 分块计算的核心手段。3. tensor 前 n 个数据计算样例 —— count 模式asc.fused_mul_add(dst, src0, src1, count512)count 模式最为简洁直接指定参与计算的元素个数为 512接口内部自动完成 mask 与迭代划分无需手动构造BinaryRepeatParams。完整的可运行核函数骨架参考单元测试 test_vector_binary.py一个同时覆盖三种形态的最小核函数如下import asc asc.jit def fused_mul_add_kernel(): x_local asc.LocalTensor(dtypeasc.float16, posasc.TPosition.VECIN, addr0, tile_size512) y_local asc.LocalTensor(dtypeasc.float16, posasc.TPosition.VECIN, addr0, tile_size512) z_local asc.LocalTensor(dtypeasc.float16, posasc.TPosition.VECOUT, addr0, tile_size512) # count 模式 asc.fused_mul_add(z_local, x_local, y_local, count512) # mask 整数连续模式 params asc.BinaryRepeatParams(1, 1, 1, 8, 8, 8) asc.fused_mul_add(z_local, x_local, y_local, mask512, repeat_times1, repeat_paramsparams) # mask 列表逐 bit模式 uint64_max 2**64 - 1 mask [uint64_max, uint64_max] asc.fused_mul_add(z_local, x_local, y_local, maskmask, repeat_times1, repeat_paramsparams) fused_mul_add_kernel[1]()源码级实现原理从 Python 调用到 MLIR 指令fused_mul_add的 Python 前端实现位于 vec_binary.pyrequire_jit set_binary_docstring(cpp_nameFusedMulAdd, append_text按元素将src0和dst相乘并加上src1最终结果存放入dst。) def fused_mul_add(dst: LocalTensor, src0: LocalTensor, src1: LocalTensor, *args, **kwargs) - None: builder global_builder.get_ir_builder() op_impl(fused_mul_add, dst, src0, src1, args, kwargs, builder.create_asc_FusedMulAddL0Op, builder.create_asc_FusedMulAddL1Op, builder.create_asc_FusedMulAddL2Op)其底层分发逻辑集中在 utils.py 的op_impl中整个链路可以概括为四步类型校验check_type依据算子名查表校验 dst/src0/src1 的 dtype 与一致性不满足即抛TypeError。重载分发OverloadDispatcher按实参特征匹配三种注册形态——mask: RuntimeIntrepeat_timesrepeat_params→ 构造FusedMulAddL0Opmask: list→ 将列表中每个元素materialize_ir_value为 uint64 IR 值后构造FusedMulAddL1Opcount: RuntimeInt→ 将 count 物化为 int32 后构造FusedMulAddL2Op。IR 构建dst.to_ir()、src0.to_ir()、src1.to_ir()将 LocalTensor 转为 IR 句柄mask/repeat_times/count 分别按 int64/int8/int32 物化repeat_params.to_ir()展开为BinaryRepeatParamsType构造值最终生成 MLIR 算子。后续编译流水MLIR 算子经 lib/Dialect/Asc 的方言处理与 lib/Target/AscendC 的代码生成最终落为 Ascend C 源码与可执行指令。在 MLIR 方言层面fused_mul_add与add、max、min、sub_relu等共用同一套 TableGen 模板族定义于 include/ascir/Dialect/Asc/IR/Basic/OpVecBinary.tddefm FusedMulAdd : BinaryTemplateL012Opfused_mul_add, FusedMulAdd; defm FusedMulAddRelu : BinaryTemplateL012Opfused_mul_add_relu, FusedMulAddRelu;BinaryTemplateL012Op模板族一次性生成了 L0 / L1 / L2 三个层级的算子定义与 Python 端三个 builder 一一对应L0 对应 mask 整数连续形态、L1 对应 mask 列表逐 bit形态、L2 对应 count 形态这也解释了为什么同一个 Python 接口在不同参数形态下会落到不同硬件指令。关联算子与扩展阅读fused_mul_add_relu在fused_mul_add的基础上追加 ReLU 激活语义为按元素将 src0 和 dst 相乘并加上 src1再进行 Relu 计算结果和 0 对比取较大值最终结果存放进 dstvec_binary.py。它同样只支持float16/float32并共享 L0/L1/L2 三级分发链路。mul_add_dst另一类乘加融合接口对应不同的操作数组织方式可对比学习向量流水线对操作数角色的区分。接口索引完整接口列表见 docs/python-api/language/basic.md 与 RST 索引 docs/python-api/rst/language/basic.rstLocalTensor的创建与属性参见 docs/python-api/language/core.md。仓库示例向量类算子更完整的工程化写法可参考 examples/01_add/add.py 与 examples/05_matmul_leakyrelu/matmul_leakyrelu.py 等示例目录。常见问题与使用建议选 count 还是 mask 形态对整段连续、无空洞的数据优先使用 count 模式——接口自动处理 mask 与迭代划分代码最简、最不易出错需要高维切分、非连续访问或精细控制每次迭代计算范围时使用 mask repeat_timesBinaryRepeatParams组合。mask 数值的含义mask 整数模式下掩码值直接对应 bit 位如mask128表示 128 个 bit 置位mask 列表模式下每个 uint64 元素控制一次迭代内的 64 个元素列表长度与迭代内 block 数匹配。步长参数何时需要修改默认blk_stride1, rep_stride8适用于数据连续排布仅在 tensor 存在行/列间隔如高维 slice、非紧密布局时才需要按 32B block 单位计算实际步长。类型检查fused_mul_add仅支持float16与float32且 dst/src0/src1 三者必须同类型混合整型会直接触发TypeError。需要整型乘加时应使用add、mul等支持int16/int32的接口。内存约束务必遵守地址对齐与地址重叠约束dst 既是目的又是乘法操作数规划地址时需避免未定义的读写冲突。【免费下载链接】pyasc本项目为Python用户提供算子编程接口支持在昇腾AI处理器上加速计算接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表