CANN PyAsc 中 asc.language.basic.scatter 详解:按偏移地址将本地张量数据分散写入 dst 的 API 用法与源码实现
2026/9/18 19:04:49 网站建设 项目流程

CANN PyAsc 中 asc.language.basic.scatter 详解:按偏移地址将本地张量数据分散写入 dst 的 API 用法与源码实现

【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc

asc.language.basic.scatter是 PyAsc 基础库中用于"按地址偏移分散数据"的向量 API,它对应 Ascend C 的Scatter函数,将源操作数src中的元素按照dst_offsetdst_base共同指定的位置写入目的操作数dst。本文基于 API 文档 与仓库源码,讲清三个函数重载的语义、各参数的单位与对齐要求、Python 侧的多路分派机制,以及 MLIR Op 到 Ascend C 代码的最终发射链路,帮助你在编写昇腾向量算子时正确使用 scatter 完成数据重排。

指令语义与对应 Ascend C 原型

scatter 指令的语义是:给定一个连续的输入张量(src)和一个目的地址偏移张量(dst_offset),根据偏移地址生成新的结果张量,并将输入张量分散到结果张量中;即把src的每个元素按照指定位置写入dst

在 PyAsc 中,该指令由 scatter 函数提供,通过asc.language.basic包对外导出(见init.py 中的from .vec_scatter import scatter__all__条目)。它与 Ascend C 的三个Scatter模板原型一一对应:

// 基于 count 的重载:处理 count 个元素 template <typename T> __aicore__ inline void Scatter(const LocalTensor<T>& dst, const LocalTensor<T>& src, const LocalTensor<uint32_t>& dstOffset, const uint32_t dstBaseAddr, const uint32_t count)
// mask 逐 bit 模式:mask 为 uint64_t 数组 template <typename T> __aicore__ inline void Scatter(const LocalTensor<T>& dst, const LocalTensor<T>& src, const LocalTensor<uint32_t>& dstOffset, const uint32_t dstBaseAddr, const uint64_t mask[], const uint8_t repeatTime, const uint8_t srcRepStride)
// mask 连续模式:mask 为单个 uint64_t template <typename T> __aicore__ inline void Scatter(const LocalTensor<T>& dst, const LocalTensor<T>& src, const LocalTensor<uint32_t>& dstOffset, const uint32_t dstBaseAddr, const uint64_t mask, const uint8_t repeatTime, const uint8_t srcRepStride)

也就是说,PyAsc 侧的三个 Python 重载分别映射到这三个 C++ 原型:count版本对应第一式,mask: int版本对应第三式(连续模式),mask: List[int]版本对应第二式(逐 bit 模式,mask[]数组)。

三个函数重载与参数说明

文档定义了asc.language.basic.scatter的三个重载签名:

# 重载 1:mask 为 int(连续模式) asc.language.basic.scatter(dst: LocalTensor, src: LocalTensor, dst_offset: LocalTensor, dst_base: int, mask: int, repeat_times: int, src_rep_stride: int) -> None # 重载 2:mask 为 List[int](逐 bit 模式) asc.language.basic.scatter(dst: LocalTensor, src: LocalTensor, dst_offset: LocalTensor, dst_base: int, mask: List[int], repeat_times: int, src_rep_stride: int) -> None # 重载 3:基于 count asc.language.basic.scatter(dst: LocalTensor, src: LocalTensor, dst_offset: LocalTensor, dst_base: int, count: int) -> None

各参数含义如下(以文档说明为准,并结合源码补充了类型约束):

参数类型说明
dstLocalTensor目的操作数,元素被写入的本地张量。
srcLocalTensor源操作数,数据类型需与dst保持一致。
dst_offsetLocalTensoruint32存储src每个元素在dst中对应的地址偏移,以字节为单位;偏移基于dst的基地址dst_base计算,取值应保证按dst数据类型位宽对齐。源码中对 dtype 有显式校验,见下文。
dst_baseintdst的起始偏移地址,单位字节,取值应保证按dst数据类型位宽对齐。
countint执行处理的数据个数(仅重载 3)。
maskintList[int]控制每次迭代内参与计算的元素,支持连续模式(单个int)或逐 bit 模式(List[int])。
repeat_timesint指令迭代次数,每次迭代完成 8 个 datablock 的数据收集。
src_rep_strideint相邻迭代间的地址步长,单位是 datablock。

其中 datablock 是 Ascend AI 处理器向量计算的基本数据单元(每 core 每 cycle 处理的 32 字节单位,即 512bit),repeat_timessrc_rep_stride共同决定了多轮迭代时src数据的读取范围。

源码实现:Python 侧的多路分派与类型校验

阅读 vec_scatter.py 可以看到,scatter的 Python 实现采用"运行时参数分派"(OverloadDispatcher)来区分三个重载:

def op_impl(callee, dst, src, dst_offset, dst_base, args, kwargs, build_l0, build_l1, build_l2) -> None: builder = build_l0.__self__ dispatcher = OverloadDispatcher(callee) check_type(dst_offset) # 重载 1:mask 为 RuntimeInt(连续模式) @dispatcher.register_auto def _(mask: RuntimeInt, repeat_times: RuntimeInt, src_rep_stride: RuntimeInt): build_l0(dst.to_ir(), src.to_ir(), dst_offset.to_ir(), _mat(dst_base, KT.uint32).to_ir(), _mat(mask, KT.uint64).to_ir(), _mat(repeat_times, KT.uint8).to_ir(), _mat(src_rep_stride, KT.uint8).to_ir()) # 重载 2:mask 为 list(逐 bit 模式) @dispatcher.register_auto def _(mask: list, repeat_times: RuntimeInt, src_rep_stride: RuntimeInt): mask = [_mat(v, KT.uint64).to_ir() for v in mask] build_l1(...) # 重载 3:count @dispatcher.register_auto def _(count: RuntimeInt): build_l2(...)

这里有两个值得注意的实现细节:

  1. dst_offset的 dtype 强校验check_type要求dst_offset必须是uint32类型,否则抛出TypeError(见 check_type):

    def check_type(dst_offset: LocalTensor) -> None: if dst_offset.dtype != KT.uint32: raise TypeError(f"Invalid dst_offset data type, got {dst_offset.dtype}, expect uint32.")

    这与 Ascend C 原型中LocalTensor<uint32_t> dstOffset的约束一致,说明 Python 侧必须在构造偏移张量时就保证类型正确,而不是等到编译期。

  2. 标量参数的类型物化dst_basemaskrepeat_timessrc_rep_stridecount等 Python 整数通过_matmaterialize_ir_value)转换为对应 IR 类型:dst_base物化为uint32mask物化为uint64repeat_timessrc_rep_stride物化为uint8count物化为uint32。这与 C++ 原型中uint32_t dstBaseAddruint64_t maskuint8_t repeatTimeuint8_t srcRepStrideuint32_t count的位宽完全对应。逐 bit 模式下列表中的每个元素都物化为uint64,与const uint64_t mask[]一致。

函数入口scatter上带有@require_jit装饰器,意味着它必须在 JIT 编译上下文(即算子 kernel 的 tracing 环境)中调用,调用时会通过global_builder.get_ir_builder()获取当前 IR builder 并创建对应的scatter_l0/l1/l2Op。

IR 层定义与 Ascend C 代码发射链路

Python API 创建的 Op 在 MLIR 层的定义位于 OpVecScatter.td,三个 Op 一一对应:

  • AscendC_ScatterL0Op(op namescatter_l0):单标量mask版本;
  • AscendC_ScatterL1Op(op namescatter_l1):Variadic<AnyType>:$mask,即 mask 为可变长参数(对应逐 bit 模式的数组);
  • AscendC_ScatterL2Op(op namescatter_l2):count版本。

注意ScatterL1Opmask是变长参数,这正是逐 bit 模式需要传入数组在 IR 层的体现。

代码发射(Emit)阶段,普通重载直接按位置参数打印出AscendC::Scatter(dst, src, dstOffset, dstBase, mask, repeatTimes, srcRepStride);而ScatterL1Op有专门的 printOperation 处理:它先把变长的 mask 参数落地为一个局部的uint64_t数组变量(printMask生成数组名),再调用带数组参数的Scatter重载——因为 Ascend C 中逐 bit 模式要求传入数组实参。这一点可以从 vec_scatter.mlir 测试用例的期望输出得到确认:

// CHECK-LABEL: void emit_scatter(AscendC::LocalTensor<float> v1, ..., uint64_t v9, uint64_t v10) { // CHECK-NEXT: AscendC::Scatter(v1, v2, v3, v4, v5, v6, v7); // CHECK-NEXT: uint64_t v1_mask_list0[] = {v9, v10}; // CHECK-NEXT: AscendC::Scatter(v1, v2, v3, v4, v1_mask_list0, v6, v7); // CHECK-NEXT: AscendC::Scatter(v1, v2, v3, v4, v8); // CHECK-NEXT: return; // CHECK-NEXT: } func.func @emit_scatter(%dst: !ascendc.local_tensor<1024xf32>, ...) { ascendc.scatter_l0 %dst, %src, %dstOffset, %dstBase, %mask, %repeatTimes, %srcRepStride : ... ascendc.scatter_l1 %dst, %src, %dstOffset, %dstBase, %maskArray1_0, %maskArray1_1, %repeatTimes, %srcRepStride : ... ascendc.scatter_l2 %dst, %src, %dstOffset, %dstBase, %count : ... return }

scatter_l1在发射时先生成uint64_t v1_mask_list0[] = {v9, v10};,再以其作为数组实参调用Scatter,完整还原了 Ascend C 逐 bit 模式的调用形态。

调用示例:三种典型场景

文档给出了三类典型用法,以下结合示例说明其适用场景。

场景一:tensor 高维切分计算——mask 连续模式

当按固定块(datablock)粒度处理高维切分的数据时,使用连续 mask:

asc.scatter(dst, src, dst_offset, dst_base=0, mask=128, repeat_times=1, src_rep_stride=8)

其中mask=128(连续模式下控制每次迭代参与计算的元素模式)、repeat_times=1表示只迭代一次(一次迭代覆盖 8 个 datablock)、src_rep_stride=8表示相邻迭代间按 8 个 datablock 步进。

场景二:tensor 高维切分计算——mask 逐 bit 模式

当需要对迭代内每个元素做细粒度(逐 bit)使能控制时,传入 mask 列表:

mask_bits = [uint64_max, uint64_max] asc.scatter(dst, src, dst_offset, dst_base=0, mask=mask_bits, repeat_times=1, src_rep_stride=8)

每个uint64的每一位对应迭代窗口内一个元素的参与开关;uint64_max表示全部位有效。逐 bit 模式在 IR 上对应scatter_l1,发射时会自动生成uint64_tmask 数组(见上文发射链路说明)。

场景三:处理源张量的前 n 个数据——count 模式

当只需处理src的前count个元素(例如源操作数实际是标量或前缀数据)时:

asc.scatter(dst, src, dst_offset, dst_base=0, count=128)

此时不使用 mask/迭代参数,而是直接指定参与处理的数据个数为 128。

仓库中的单元测试 test_scatter 在同一个 kernel 中依次覆盖了这三种调用形态,可作为最小可运行的参考骨架:

def kernel_scatter() -> None: ... asc.scatter(dst, src, dst_offset, dst_base=0, count=128) asc.scatter(dst, src, dst_offset, dst_base=0, mask=128, repeat_times=1, src_rep_stride=8) mask_bits = [uint64_max, uint64_max] asc.scatter(dst, src, dst_offset, dst_base=0, mask=mask_bits, repeat_times=1, src_rep_stride=8)

使用建议与注意事项

结合文档约束与源码校验逻辑,实际使用时建议关注以下几点:

  1. 偏移对齐dst_offsetdst_base均以字节为单位,且必须按dst的数据类型位宽对齐。例如dstfloat16(2 字节)时,偏移量应保证为 2 的整数倍;float32时保证为 4 的整数倍。
  2. 类型匹配srcdst的数据类型需一致;dst_offset必须为uint32张量,否则 Python 侧会立即抛出TypeError
  3. 迭代窗口计算:每次迭代固定覆盖 8 个 datablock,repeat_times决定迭代轮数,src_rep_stride(datablock 单位)决定轮间步长。设计 mask 时应据此计算总处理规模,避免越界。
  4. JIT 上下文scatter@require_jit约束,只能在算子 kernel 的 JIT tracing 上下文中调用,不能脱离 builder 环境单独执行。
  5. 与 gather 的对称性:scatter 与 gather 一样都是按偏移重排数据的本地内存指令,二者常成对出现于需要索引搬运的算子中;scatter 侧重"按偏移写入",方向与 gather 相反。

小结

asc.language.basic.scatter是 PyAsc 中把"按偏移地址分散写入"这一 Ascend C 能力 Python 化的接口:三个重载分别对应count版本、mask 连续模式和 mask 逐 bit 模式,参数单位(字节 / datablock)与对齐要求需在编写算子时严格遵守。其实现链路清晰可溯——Python 侧 vec_scatter.py 完成类型校验与参数物化,MLIR 侧 OpVecScatter.td 定义scatter_l0/l1/l2三个 Op,VecScatter.cpp 负责把逐 bit 模式的变长 mask 落地为uint64_t数组并还原为 Ascend C 的Scatter调用,测试用例 vec_scatter.mlir 与 test_common_api.py 则验证了从 IR 到 C++ 代码的完整一致性。

【免费下载链接】pyasc本项目为Python用户提供算子编程接口,支持在昇腾AI处理器上加速计算,接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyasc

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询