PyPTO 中 split 算子的 view 切片实现:batch 轴 loop 搬运与轴切分骨架解析
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
导读
本文基于 CANN pypto-gym 仓库中 pypto-api-explore 技能的 split kernel 参考骨架,系统讲解如何在 PyPTO 编程框架下实现split类轴切分算子:batch 轴用pypto.loop循环搬运,切分轴用pypto.view取一片(view 切片),其余轴整块搬运。读完本文,你将掌握pypto.loop/pypto.view/pypto.set_vec_tile_shapes/pypto.assemble四个接口在轴切分场景下的组合方式,理解占位符参数(sl/ol/batch/inner/inner_out)的语义,并能结合仓库中的真实算子实现(如 Arctic LSTM、Qwen3.5 GDR)把骨架扩展为可运行的 NPU kernel。
split 的 Torch→PyPTO 映射:view 切片而非独立 API
PyPTO 框架没有为split提供独立的原子接口。仓库中的 Torch ↔ Pypto 算子对标手册 明确将split/chunk归类为「差异映射」条目:
split/chunk→view(切片):需按份数计算各片偏移
也就是说,torch.split(tensor, split_size_or_sections, dim)在 PyPTO 中需要由开发者手工完成两件事:
- 按份数计算各片的起始偏移(offset),这是
split与普通view的关键差异; - 用
pypto.view从原张量上截取对应切片,切片仍是原 GM 数据的视图,不产生额外搬运。
而由于 NPU 片上存储(UB)容量有限,一个完整 shape 的切片往往无法一次放入,因此 split 骨架采用了「batch 轴 loop、逐片搬运」的经典 Tiling 策略——这正是本文 split.md 骨架的核心思想。
骨架代码全解:四步完成一次切分
split.md 给出的完整骨架如下(注释为逐行解析):
@pypto.frontend.jit(runtime_options={"run_mode": pypto.RunMode.NPU}) def split_kernel(a: pypto.Tensor(sl, pypto_dtype), out: pypto.Tensor(ol, pypto_dtype)): # 第 1 步:batch 轴 loop 搬运。i 为当前迭代的 batch 偏移, # unroll_list=[1] 表示按 1 个 batch 为单位展开/流水 for i in pypto.loop(batch, name="batch", unroll_list=[1]): # 第 2 步:从输入 a 上取"当前 batch 行 + 全部内层"的切片 a_s。 # 偏移量 = [i, 0, 0, ...],即 batch 轴偏移 i,其余轴偏移 0(整块) a_s = pypto.view(a, [1] + inner, [i] + [0] * len(inner)) # 第 3 步:声明本次迭代的向量 tile shape(UB 分块尺寸)。 # 首维为 1(当前 batch),其余维度为内层输出形状 pypto.set_vec_tile_shapes(1, *inner_out) # 第 4 步:在 a_s 内再取一片 r:shape 为 [1] + inner_out, # 偏移全部为 0(因为 a_s 已经是当前 batch 行) r = pypto.view(a_s, [1] + inner_out, [0] * len([1] + inner)) # 第 5 步:把 r 装配(写回)到输出 out 的对应位置: # 输出偏移 = [i, 0, 0, ...],即 batch 轴写回第 i 行 pypto.assemble(r, [i] + [0] * (len(ol) - 1), out)逐接口语义拆解
pypto.loop(batch, name="batch", unroll_list=[1]):按batch长度建立循环。name用于在编译产物/Profile 中标识循环;unroll_list声明允许的展开粒度(当前骨架用[1],即一次处理 1 个 batch 行)。仓库中更复杂的实现会使用带偏移与展开长度的pypto.loop_unroll,例如 sum_lstm.py 中:
for bs_offset, unroll_length in pypto.loop_unroll( 0, batch_size, 1, name="LSTM_BATCH_LOOP", idx_name="bs_offset", unroll_list=tile_config.unroll_list ):可见unroll_list是可调优参数——sum_lstm.py 中 Arctic LSTM 将其配置为[1, 2, 4],以配合不同的 batch tile 大小。
pypto.view(a, shape, offsets):从 GM 张量上按shape与offsets截取一个 view 切片,不拷贝数据、不产生 GM 搬运。注意 offsets 列表的长度必须等于张量维数,每个元素表示对应维度的起始下标。骨架第 2 步中[i] + [0] * len(inner)表示「batch 轴偏移 i、其余轴从 0 开始整块取」。
pypto.set_vec_tile_shapes(1, *inner_out):声明本次迭代的 Vector 单元 tile 形状(UB 分块)。首维固定 1(单 batch 行),后续维度来自inner_out。tile shape 决定 UB 装载粒度,需按实际 shape / dtype / 平台约束调优,而非照抄骨架。值得注意:一个 kernel 内可以多次调用它来切换不同阶段的 tile 形状,例如 sum_lstm.py 先用(current_tile_bs, h_tile),随后在门控阶段又切回(1, hidden_dim_4)。
pypto.assemble(r, offsets, out):把片上计算结果r写回 GM 输出out的offsets位置。这里偏移[i] + [0] * (len(ol) - 1)把第 i 行写回输出第 i 行,其余轴从 0 开始——与pypto.view的读偏移形成「读一片 → 算一片 → 写一片」的对称结构。
切分策略一句话总结
骨架 Note 原文为:
batch 轴 loop 搬运;沿切分轴取一片(view 切片),其余轴整块。
这正是 PyPTO 轴切分算子的通用模式:loop 覆盖 batch(GM 搬运的单位),view 负责在切分轴上精确定位切片,其余轴整块随行搬运。
占位符约定与最小可运行 setup
骨架使用了一组占位符,语义约定见 examples/README.md:
| 占位符 | 含义 |
|---|---|
sl | 输入 shape 列表,如[B, S, D] |
ol | 输出 shape 列表 |
pypto_dtype | 元素 dtype,如pypto.DT_FP32 |
batch | 被 loop 的外层轴长度(通常sl[0]) |
inner | 单次迭代处理的内层 shape,如sl[1:] |
inner_out | 单次迭代的输出内层 shape(末轴长度可能变化,如 glu/diff/mean) |
配套的最小可运行 setup(同一个 README 提供):
import pypto B, D = 8, 128 sl, ol = [B, D], [B, 1] pypto_dtype = pypto.DT_FP32 batch, inner, inner_out = B, [D], [1]代入骨架后,语义即为:输入[8, 128]的 FP32 张量,batch 轴8逐行 loop,每行view出[1, 128]切片,tile shape 为(1, 1),最终装配到输出[8, 1]。注意inner_out的末轴长度可以和输入不同(如 split 出子序列、或经 glu/diff 等变换后维度变化),这正是占位符拆分inner/inner_out的原因。
真实源码佐证:view + assemble 的实战写法
骨架是「参考形态」,仓库中的生产级实现展示了更丰富的用法:
1. 多片切分(同一输入多个 offsets):sum_lstm.py 将融合后的[BS, 4H]张量按 4 等分切出 LSTM 四门:
pre_f = pypto.view(fused, [current_tile_bs, hidden_dim], [0, 0]) pre_i = pypto.view(fused, [current_tile_bs, hidden_dim], [0, hidden_dim * 1]) pre_o = pypto.view(fused, [current_tile_bs, hidden_dim], [0, hidden_dim * 2]) pre_c = pypto.view(fused, [current_tile_bs, hidden_dim], [0, hidden_dim * 3])这正是 split 的典型应用:切分轴偏移 = 片序号 × 每片长度,其余轴偏移 0。多个 view 共享同一份 GM 数据,零拷贝切出多片。
2. 带 valid_shape 的尾块处理:gdr_bwd_impl.py 展示了 view 的第 4 个参数valid_shape,用于动态长度(如序列尾块act < bt)时声明有效形状:
pypto.view(q, [bt, head_dim], [off, hc], valid_shape=[act, head_dim]), ... b_f = pypto.fillpad(pypto.view(beta, [bt, 1], [off, h1], valid_shape=[act, 1]), "constant", 0.0)当 split 的切分块长度不能整除原轴长时(尾片不足),可组合valid_shape与pypto.fillpad做补零对齐(该组合在 Qwen3-Next gated_delta_rule、chunk_kda 中用于 q/k/v/beta 尾块填充,见 pypto-specific-ops.md 中fillpad条目)。
3. assemble 的多种写回位置:gdr_bwd_impl.py 展示了assemble同时支持「view 切片后写回」与「整块直接写回」两种形态,并支持 per-head 偏移[off, hc]累加写回。
与 chunk 骨架的关系:同一模式的不同切分粒度
chunk.md 与 split.md 的骨架完全同构(loop + view + set_vec_tile_shapes + assemble 四段式),差异仅在语义层:
split(split_size):按固定长度切分,每片长度相同(除最后一片可能不足);chunk(chunks):按固定份数切分,每片长度 =ceil(轴长 / 份数)。
无论哪种语义,落到 PyPTO 上都是「先算各片 offsets,再逐片 view + assemble」,这正是映射手册将二者归为同一条目(split/chunk→view)的原因。实际开发中,每片对应的inner/inner_out与 tile shape 可能因片长不同而需要分别配置。
使用约束与调优注意点
骨架未经 NPU 编译验证:examples/README.md 明确声明「骨架未逐一经 NPU 编译验证」,且仅作「接口组合与轴切分模式」参考,不是标准模板。loop 轴、
unroll_list、tile shape、动态轴处理都需按实际 shape / dtype 与平台约束确定并调优。动态 shape 约束:按 SKILL.md 的硬约束速查,matmul / 归约类等计算 API 在编译期需要 concrete shape,不接受含
DYNAMIC维度的 tensor。若 split 算子中 batch 轴为动态(如 sum_lstm.py 中的pypto.DYNAMIC),应保证只有 loop 轴是动态的、tile 内维度为静态(pypto.STATIC),否则需采用「loop 切 tile」策略规避。tile shape 与 UB 容量:
set_vec_tile_shapes的取值需保证单次迭代数据能放入 UB,且每维 > 0、最多 4 维;不匹配会导致编译失败或性能劣化。可参考execution-constraints.md与 Tiling 文档(pypto-set_vec_tile_shapes.md)核对约束。view 是切片不是拷贝:切出的 view 与原张量共享底层 GM 数据,因此读侧不会产生额外搬运;而写回必须通过
assemble显式完成,且输出 offsets 必须与循环变量对齐,避免多核/多迭代写同一位置产生数据竞争。多片输出扩展:骨架只展示了切出一片(单输出)的形态;当需要把输入切成 N 片写到 N 个输出时,在循环体内对每片重复「view → 计算 → assemble」即可,各片 offsets 按
片号 × 片长递增计算。
小结
PyPTO 中实现split的标准姿势可以概括为一条公式:batch轴pypto.loop搬运 +pypto.view按计算好的 offsets 取切片 +pypto.set_vec_tile_shapes声明 UB 分块 +pypto.assemble写回对应输出位置。仓库提供的 split.md 骨架是这一模式的最小化抽象,而 sum_lstm.py(多片切分)、gdr_bwd_impl.py(动态 valid_shape 尾块)则展示了其在真实 NPU 算子中的完整演进形态——理解这四步,即可举一反三地编写任意的轴切分类 kernel。
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考