ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

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

CANN PyAsc 中 asc.language.basic.scatter 详解:按偏移地址将本地张量数据分散写入 dst 的 API 用法与源码实现 CANN PyAsc 中 asc.language.basic.scatter 详解按偏移地址将本地张量数据分散写入 dst 的 API 用法与源码实现【免费下载链接】pyasc本项目为Python用户提供算子编程接口支持在昇腾AI处理器上加速计算接口与Ascend C一一对应并遵守Python原生语法。项目地址: https://gitcode.com/cann/pyascasc.language.basic.scatter是 PyAsc 基础库中用于按地址偏移分散数据的向量 API它对应 Ascend C 的Scatter函数将源操作数src中的元素按照dst_offset与dst_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 LocalTensorT dst, const LocalTensorT src, const LocalTensoruint32_t dstOffset, const uint32_t dstBaseAddr, const uint32_t count)// mask 逐 bit 模式mask 为 uint64_t 数组 template typename T __aicore__ inline void Scatter(const LocalTensorT dst, const LocalTensorT src, const LocalTensoruint32_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 LocalTensorT dst, const LocalTensorT src, const LocalTensoruint32_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的三个重载签名# 重载 1mask 为 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 # 重载 2mask 为 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。maskint或List[int]控制每次迭代内参与计算的元素支持连续模式单个int或逐 bit 模式List[int]。repeat_timesint指令迭代次数每次迭代完成 8 个 datablock 的数据收集。src_rep_strideint相邻迭代间的地址步长单位是 datablock。其中 datablock 是 Ascend AI 处理器向量计算的基本数据单元每 core 每 cycle 处理的 32 字节单位即 512bitrepeat_times与src_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) # 重载 1mask 为 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()) # 重载 2mask 为 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(...) # 重载 3count dispatcher.register_auto def _(count: RuntimeInt): build_l2(...)这里有两个值得注意的实现细节dst_offset的 dtype 强校验。check_type要求dst_offset必须是uint32类型否则抛出TypeError见 check_typedef check_type(dst_offset: LocalTensor) - None: if dst_offset.dtype ! KT.uint32: raise TypeError(fInvalid dst_offset data type, got {dst_offset.dtype}, expect uint32.)这与 Ascend C 原型中LocalTensoruint32_t dstOffset的约束一致说明 Python 侧必须在构造偏移张量时就保证类型正确而不是等到编译期。标量参数的类型物化。dst_base、mask、repeat_times、src_rep_stride、count等 Python 整数通过_matmaterialize_ir_value转换为对应 IR 类型dst_base物化为uint32mask物化为uint64repeat_times与src_rep_stride物化为uint8count物化为uint32。这与 C 原型中uint32_t dstBaseAddr、uint64_t mask、uint8_t repeatTime、uint8_t srcRepStride、uint32_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_ScatterL0Opop namescatter_l0单标量mask版本AscendC_ScatterL1Opop namescatter_l1VariadicAnyType:$mask即 mask 为可变长参数对应逐 bit 模式的数组AscendC_ScatterL2Opop namescatter_l2count版本。注意ScatterL1Op的mask是变长参数这正是逐 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::LocalTensorfloat 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_tensor1024xf32, ...) { 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粒度处理高维切分的数据时使用连续 maskasc.scatter(dst, src, dst_offset, dst_base0, mask128, repeat_times1, src_rep_stride8)其中mask128连续模式下控制每次迭代参与计算的元素模式、repeat_times1表示只迭代一次一次迭代覆盖 8 个 datablock、src_rep_stride8表示相邻迭代间按 8 个 datablock 步进。场景二tensor 高维切分计算——mask 逐 bit 模式当需要对迭代内每个元素做细粒度逐 bit使能控制时传入 mask 列表mask_bits [uint64_max, uint64_max] asc.scatter(dst, src, dst_offset, dst_base0, maskmask_bits, repeat_times1, src_rep_stride8)每个uint64的每一位对应迭代窗口内一个元素的参与开关uint64_max表示全部位有效。逐 bit 模式在 IR 上对应scatter_l1发射时会自动生成uint64_tmask 数组见上文发射链路说明。场景三处理源张量的前 n 个数据——count 模式当只需处理src的前count个元素例如源操作数实际是标量或前缀数据时asc.scatter(dst, src, dst_offset, dst_base0, count128)此时不使用 mask/迭代参数而是直接指定参与处理的数据个数为 128。仓库中的单元测试 test_scatter 在同一个 kernel 中依次覆盖了这三种调用形态可作为最小可运行的参考骨架def kernel_scatter() - None: ... asc.scatter(dst, src, dst_offset, dst_base0, count128) asc.scatter(dst, src, dst_offset, dst_base0, mask128, repeat_times1, src_rep_stride8) mask_bits [uint64_max, uint64_max] asc.scatter(dst, src, dst_offset, dst_base0, maskmask_bits, repeat_times1, src_rep_stride8)使用建议与注意事项结合文档约束与源码校验逻辑实际使用时建议关注以下几点偏移对齐dst_offset与dst_base均以字节为单位且必须按dst的数据类型位宽对齐。例如dst为float162 字节时偏移量应保证为 2 的整数倍float32时保证为 4 的整数倍。类型匹配src与dst的数据类型需一致dst_offset必须为uint32张量否则 Python 侧会立即抛出TypeError。迭代窗口计算每次迭代固定覆盖 8 个 datablockrepeat_times决定迭代轮数src_rep_stridedatablock 单位决定轮间步长。设计 mask 时应据此计算总处理规模避免越界。JIT 上下文scatter带require_jit约束只能在算子 kernel 的 JIT tracing 上下文中调用不能脱离 builder 环境单独执行。与 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三个 OpVecScatter.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),仅供参考
返回列表