PyPTO 数学函数 expands 详解:Tile 标量填充(splat)的原理与实战用法
2026/9/20 2:44:29 网站建设 项目流程
  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

PyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。

项目地址:https://gitcode.com/cann/pypto
点击查看免费下载

导读

pypto_pro.language.expands是 CANN PyPTO 编程范式中用于将整块 Tile 的所有元素填充为同一个标量值的数学函数,其语义等价于out[i] = scalar,在 Kernel 编程中常用于初始化负无穷 Tile(构造因果掩码)或零 Tile(累加器清零)。本文以 expands.md 为骨架,完整梳理该函数的函数原型、参数约束、支持的数据类型与运行平台,并结合仓库源码(Python API 声明、IR 构建、类型与取值范围校验、ST/UT 测试)深入讲解其底层实现原理,最终给出可直接运行的基本用法与典型业务场景示例,帮助读者在 PyPTO 矢量编程中正确、高效地使用标量填充操作。

功能说明

expands用于将目的 Tile 填充为指定标量值,其计算语义为:

out[i] = scalar (对 out 中的每一个元素 i)

在 PyPTO 的 Tile 编程模型中,expands属于**矢量计算(Vector/SIMD)**操作,需要在section_vector()矢量流水段中执行。该函数最常见的两个用途是:

  • 初始化负无穷 Tile:为 Softmax、注意力(Attention)等算子构造因果掩码时,需要将掩码区域填充为-inf,保证指数运算后该区域权重归零;
  • 初始化零 Tile:为归约、累加等算子预先清零,或为不同数据类型的 Tile 提供统一的零值初始化入口。

从源码注释看,其本质是一次"标量广播"(splat)操作:API 声明 中明确写为"""Fill Tile with a scalar (splat): ``out[i] = scalar``"""

产品支持情况

expands的硬件支持情况与 PyPTO 的矢量流水线支持范围一致,具体如下:

产品支持情况
Ascend 950PR / Ascend 950DT支持
Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持
Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持

说明:以上支持矩阵来自 expands.md 的"产品支持情况"章节。在调用前请确认目标运行环境属于 Ascend 950 系列,否则编译或运行阶段可能报错。

函数原型

pypto_pro.language.expands( out: Tile, scalar: Scalar, ) -> None

函数位于pypto_pro.language命名空间下,通过import pypto_pro.language as pl后以pl.expands(...)方式调用。它是一个无返回值的原地写操作:结果直接写入out指向的 Tile,函数返回None

参数说明

参数输入/输出说明
out输出目的操作数,Tile类型,全部元素被填充为scalar值。
数据类型支持:DT_UINT8DT_INT8DT_UINT16DT_INT16DT_UINT32DT_INT32DT_INT64DT_UINT64DT_FP16DT_BF16DT_FP32
位于 UB(Unified Buffer)或 L1 Buffer。
scalar输入填充值。为整型或浮点型常量,或运行时整型或浮点型标量表达式,类型须与out元素类型兼容。

关键点解读

  • out的数据类型覆盖整型(UINT8 至 UINT64)与浮点型(FP16/BF16/FP32)共 11 种,未包含 FP64 与 FP8 系列,这与 IR 层的 dtype 白名单校验 中_EXPANDS_DTYPES的约束一一对应;
  • out必须位于 UB 或 L1 Bufferexpands是矢量单元执行的写操作,操作数必须落在矢量单元可直接访问的存储空间。在示例中通过TileType(target_memory=pl.MemorySpace.Vec)显式指定;
  • scalar的兼容性:既可以是编译期字面量常量,也可以是运行期的标量表达式(如循环变量、symbolic_scalar计算结果)。若字面量超出out元素类型的可表示范围,编译期会抛出范围校验错误(见下文"范围校验")。

约束说明

官方文档声明该函数无额外约束。需要补充说明的是,"无约束"指函数本身没有形状、步长等特殊限制,但实际使用中仍需满足 PyPTO 通用编程约束:

  • expands属于矢量操作,必须在pl.section_vector()代码块内调用;
  • 填充的 Tile 需通过pl.make_tile_group创建,并保证地址、互斥锁(mutex)分配不与同流水段其他 Tile 冲突;
  • scalar的值必须能被out的元素类型精确表示,否则在编译期触发范围检查错误。

返回值说明

无返回值(None)。填充结果直接写入out参数对应的目的 Tile。

调用示例

以下示例均可在 Ascend 950 系列设备上运行,并与仓库 ST 测试 test_reduce_expand_fp32.py 中的expands_kernel保持一致。

基本用法

import pypto_pro.language as pl K_VALUE = 2.0 @pl.jit(auto_mutex=True) def expands_kernel(dummy: pl.Tensor[[64, 64], pl.DT_FP32], out: pl.Tensor[[64, 64], pl.DT_FP32]): tt = pl.TileType(shape=[64, 64], dtype=pl.DT_FP32, target_memory=pl.MemorySpace.Vec) tile_out = pl.make_tile_group(type=tt, addrs=0x0000, mutex_ids=[0]) with pl.section_vector(): cur_out = tile_out.current() pl.expands(cur_out, K_VALUE) pl.store(out, cur_out, [0, 0])

运行结果(文档示例输出):

输入数据K_VALUE:2 输出数据out:[[2 2 2 2 2 2 2 2 ...], [2 2 2 2 2 2 2 2 ...], [2 2 2 2 2 2 2 2 ...], [2 2 2 2 2 2 2 2 ...], ...]

示例关键步骤拆解:

  1. 创建 Tile 类型TileType(shape=[64, 64], dtype=pl.DT_FP32, target_memory=pl.MemorySpace.Vec)定义了一块 64×64 的 FP32 矢量 Tile;
  2. 绑定 Tile 组make_tile_group(type=tt, addrs=0x0000, mutex_ids=[0])将 Tile 绑定到 UB 起始地址0x0000并申请互斥锁0auto_mutex=True会自动管理锁的申请与释放;
  3. 进入矢量流水段with pl.section_vector():声明后续指令在矢量单元执行;
  4. 取当前 Tile 并填充cur_out = tile_out.current()获取当前实例,pl.expands(cur_out, K_VALUE)将整块填充为2.0
  5. 写回全局内存pl.store(out, cur_out, [0, 0])将 UB 中的 Tile 存回 Tensorout

对应的 ST 测试在 test_reduce_expand_fp32.py 中验证:构造torch.full((64, 64), K_VALUE)作为参考结果,与 Kernel 输出逐元素比对(rtol=1e-2, atol=1e-2),确认整块 Tile 均被填充为2.0

初始化负无穷 Tile

# 初始化负无穷 Tile(因果掩码) pl.expands(neg_inf_vec, NEG_INF)

这是expands最具代表性的应用场景。在实现 Softmax 或 Attention 的因果掩码(causal mask)时,需要把上三角区域填充为-inf,使exp(-inf) = 0,从而在后续 softmax 归一化中屏蔽非法位置的注意力权重。与逐元素赋值相比,expands一条指令即可完成整块 Tile 的填充,避免了循环与地址计算开销。

初始化零 Tile

# 初始化零 Tile pl.expands(score_u16_row, 0)

当需要对 Tile 清零时,expands同样是最直接的方式。示例中score_u16_rowDT_UINT16类型的 Tile,填充值0与无符号整型兼容。清零操作常用于累加器初始化、输出缓冲区预填充等场景,此时expands可替代逐元素store,显著减少指令数。

源码级原理剖析

1. Python API 声明层

expands在 python/pypto_pro/language/_api.py 中声明,并通过 python/pypto_pro/language/init.py 的__all__导出,最终以pl.expands形式对用户可见。它是一个使用@_api_decl修饰的声明式 API:函数体只承载文档字符串,真正的编译逻辑由解析器与 IR 构建器接管。

2. 前端解析与操作数角色

在 PyPTO 的前端解析流水线中,expands被注册在 python/pypto_pro/language/parser/_op_pipeline.py 的操作数角色表中:

"expands": ["W", None],

该表声明了expands的两个参数中,第一个参数out写操作数(W),第二个参数scalar是非 Tile 操作数。get_op_tile_roles(python/pypto_pro/language/parser/_op_pipeline.py)在跨核同步分析等场景中依赖该角色信息,确保写操作数的跨核依赖被正确追踪。此外,python/pypto_pro/language/parser/_call_parser.py 在调用解析时对expands做了专门分支处理。

3. IR 构建与双重校验

expands的 IR 构建逻辑位于 python/pypto_pro/ir/op/block_ops.py 的_ir_expands,并注册到 block 算子表(block_ops.py 中的"expands": OpSpec(builder=_ir_expands))。其核心流程包含两道校验:

def _ir_expands(out: Expr, scalar: Expr, *, span: Span | None = None) -> Expr: out_dtype = getattr(out.type, "dtype", None) _check_dtype("expands", out_dtype, _EXPANDS_DTYPES) # The splat value lands in the out tile, so it has to be representable there. check_const_expr_fits_dtype(scalar, out_dtype, span=span, api="pl.expands") return _ir_core.create_op_call(block_ir_op("expands"), [out, scalar], {}, span or _span())
  • 第一道:dtype 白名单校验_check_dtype("expands", out_dtype, _EXPANDS_DTYPES)检查out的 dtype 是否在文档列出的 11 种类型之内,非法类型在编译期直接报错;
  • 第二道:标量取值范围校验check_const_expr_fits_dtype(scalar, out_dtype, ...)检查字面量scalar是否能被out的元素类型表示。代码注释明确指出:splat 值最终要落进 out Tile,因此必须在其 dtype 可表示范围内——如果不做此检查,标量只受 IR 存储带宽(storage band)限制,而 IR 存储带宽远宽于窄位宽的 Tile dtype(如 UINT8),可能导致溢出后静默回绕。

4. 取值范围校验的测试佐证

取值范围校验的边界行为在 UT 与 ST 测试中被完整覆盖:

  • test_scalar_range_validation.py:向DT_INT64的 Tile 填充INT64_MAX + 1(即9223372036854775808),该值虽能落入 IR 的 uint64 存储带宽,但超出 INT64 元素类型可表示范围,编译期抛出OutOfRange,错误信息为pl.expands: scalar operand must be representable in int64, i.e. in [-9223372036854775808, 9223372036854775807], got ...
  • test_scalar_range_validation.py:填充INT64_MAX(合法边界值)的 Kernel 可在 950 设备上真实运行,输出与torch.full(..., INT64_MAX, dtype=torch.int64)完全一致(rtol=0, atol=0);
  • test_scalar_range.py:UT 注释明确"pl.expands: the splat value must fit the out tile's dtype",从单测层面对该约束进行了系统验证。

这些测试共同印证了"编译期拦截非法填充值、边界合法值可正常下板执行"的实现事实。

与相近操作的对比

在 PyPTO 算子库中,expands与以下操作易混淆,使用时需注意区分:

操作语义关键差异
pl.expands将整块 Tile 填充为同一标量标量广播,无源 Tile,仅 2 个参数
pl.full以标量值填充(见 _op_pipeline.py)同属填充类算子,expands为矢量流水段内的 splat 写操作
pl.fillpad填充 Tile 的 padding 区域仅填充边缘/填充区域,非全量填充(python/pypto_pro/language/_api.py)
row_expand/col_expand将归约结果沿行/列方向扩展回原形状源是归约输出 Tile,目标是形状扩展,非标量填充(_op_pipeline.py)

需要特别说明的是,row_expand/col_expand属于**归约扩展(Reductions / expands)**类别(python/pypto_pro/language/_api.py 的 B8 分组注释),其语义是"维度广播",与expands的"标量广播"是两类不同操作,不应混淆。

小结

pypto_pro.language.expands是 PyPTO 矢量编程中最基础也最高效的标量填充原语:

  • 用法上:一条指令完成整块 Tile 的标量填充,是初始化负无穷掩码、零 Tile 的首选方案;
  • 约束上out支持 11 种常见整型/浮点 dtype,必须位于 UB 或 L1 Buffer,scalar须与out元素类型兼容;
  • 实现上:从 API 声明 到 IR 构建 再到 操作数角色表,形成"声明—解析—校验—建图"的完整链路,其中 dtype 白名单与标量可表示范围双重校验在编译期即拦截非法用法;
  • 验证上:ST 测试 test_reduce_expand_fp32.py 与范围校验测试 test_scalar_range_validation.py 共同保证了功能正确性与边界行为。

在编写需要掩码、清零或常量预填充的 PyPTO Kernel 时,优先考虑pl.expands,它能让代码更简洁,也让矢量流水线获得更高的指令效率。

  • 人工智能
  • 编译器
  • 模型编译
  • 高性能计算
  • 深度学习
  • CANN

【免费下载链接】pypto

PyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。

项目地址:https://gitcode.com/cann/pypto
点击查看免费下载

相关推荐

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

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

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

立即咨询