
PyPTO 上实现 torch.masked_fill_ 的 kernel 参考骨架gt 掩码 where 填充的 inplace 写回方案【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym本文是 PyPTO-Gym 中pypto-api-explore技能集 examples/masked_fill_inplace.md 的深度解析。它面向需要在 NPU 上以 PyPTO 编程框架实现「掩码填充」类算子的算子开发者讲解如何用gt生成掩码、where完成三路选择并通过assemble把结果写回输出缓冲以模拟 inplace 语义。读完本文你将掌握 masked_fill 类算子的轴切分思路batch 轴 loop 切分、last-dim 整块、Vector 型算子的 tile 配置以及从 torch 算子到 PyPTO 组合方案的完整映射方法。1. 背景torch.masked_fill_ 与 PyPTO 的组合实现PyTorch 的Tensor.masked_fill_(mask, value)是原地inplace语义的逐元素算子mask为True的位置用value填充其余位置保持原值结果直接写回self。在 Transformer 模型的 attention mask、position_ids 构造等场景中它被广泛用于把「被掩码位置」置为 0 或极小的负值如-1e9以屏蔽注意力。然而 PyPTO 的原子接口中不存在 inplace 语义的算子。映射手册 torch-pypto-op-mapping.md 的「索引操作」一节明确给出了组合方案Torch 算子PyPTO 组合方案参考实现masked_fillgtwheremasked_fill.mdmasked_fill_gtwheremasked_fill_inplace.md也就是说masked_fill_与masked_fill在 PyPTO 侧使用完全相同的gtwhere组合唯一的差异体现在结果如何落盘inplace 版本把结果写回out缓冲out即a的别名非 inplace 版本写入独立的输出张量。二者的完整 kernel 参考骨架见 examples/ 目录。2. kernel 参考骨架逐行拆解原文档给出了完整的 NPU 运行 kernel如下pypto.frontend.jit(runtime_options{run_mode: pypto.RunMode.NPU}) def masked_fill_inplace_kernel(a: pypto.Tensor(sl, pypto_dtype), out: pypto.Tensor(sl, pypto_dtype)): for i in pypto.loop(batch, namebatch, unroll_list[1]): a_s pypto.view(a, [1] inner, [i] [0] * len(inner)) pypto.set_vec_tile_shapes(1, *inner) zero pypto.full([1] inner, 0.0, pypto_dtype) mask pypto.gt(a_s, zero) fill_val pypto.full([1] inner, -1e9, pypto_dtype) r pypto.where(mask, fill_val, a_s) pypto.assemble(r, [i] [0] * len(inner), out)Note: batch 轴 loop 切分last-dim 整块gt 生成掩码后 where 填充out 即 inplace 写回。2.1 JIT 入口与运行模式pypto.frontend.jit(runtime_options{run_mode: pypto.RunMode.NPU})声明这是一个在 NPU 上运行的 JIT 编译 kernel。函数签名使用pypto.Tensor(sl, pypto_dtype)声明输入a与输出out其中sl是输入 shape 列表如[B, S, D]pypto_dtype是元素 dtype如pypto.DT_FP32。这里体现出 PyPTO 的一个关键设计kernel 的 I/O 都是显式缓冲区。inplace 语义不是通过「修改a」实现而是调用方将out与a指向同一块数据kernel 内把结果assemble回out从而在 NPU 上等效完成原地更新。2.2 轴切分batch 轴 looplast-dim 整块for i in pypto.loop(batch, namebatch, unroll_list[1]): a_s pypto.view(a, [1] inner, [i] [0] * len(inner))这是本骨架的核心切分策略外层 batch 轴用pypto.loop切分每次迭代通过pypto.view取出形状为[1] inner的一个切片偏移为[i] [0] * len(inner)即只在 batch 维移动而last-dim最内层计算轴整块处理不继续切 tile。unroll_list[1]表示该层循环不完全展开保留循环结构避免编译期图膨胀。这种「外层切 batch、内层整块」的模式对逐元素算子非常典型masked_fill是纯 Vector 型逐元素计算不涉及跨行归约或矩阵乘last-dim 直接整块即可获得连续访存与简单的 tile 配置。占位符batch、inner的含义与最小可运行 setup 见 examples/README.mdimport pypto B, D 8, 128 sl, ol [B, D], [B, 1] pypto_dtype pypto.DT_FP32 batch, inner, inner_out B, [D], [1]2.3 Tile 配置set_vec_tile_shapespypto.set_vec_tile_shapes(1, *inner)因为本算子只含逐元素/比较类操作按 SKILL.md 的算子类型判断规则它属于Vector 类型只需调用pypto.set_vec_tile_shapes()设置 Vector tile无需set_cube_tile_shapes()。(1, *inner)表示每个迭代切片[1] inner的 tile 形状与切片形状一致即单次迭代内不进一步分块整块送入 Vector 计算单元。需要注意 tile shape 的通用约束每维必须 0、最多 4 维。骨架中的具体值如(1, D)是占位约定实际取值应结合 shape、dtype 与平台约束在开发阶段确定。2.4 计算主体gt 生成掩码 → where 三路选择zero pypto.full([1] inner, 0.0, pypto_dtype) mask pypto.gt(a_s, zero) fill_val pypto.full([1] inner, -1e9, pypto_dtype) r pypto.where(mask, fill_val, a_s)四行代码完整实现了「掩码填充」逻辑pypto.full([1] inner, 0.0, pypto_dtype)构造与切片同形状的全 0 张量作为比较阈值pypto.gt(a_s, zero)逐元素比较a_s 0生成布尔掩码mask——这是组合方案中的gt步骤。注意这里的掩码条件是按需定义的骨架以「大于 0 才保留」为例实际算子中掩码条件如a 0、attention_mask 0等完全由业务决定只需保证比较结果为布尔张量即可喂给wherepypto.full([1] inner, -1e9, pypto_dtype)构造填充值张量。-1e9是 attention 场景中经典的「负无穷大」近似值用于在 softmax 前屏蔽被掩码位置实际使用时可按语义替换为 0 或其他值pypto.where(mask, fill_val, a_s)逐元素三路选择mask为True处取fill_val否则取a_s原值。这正是 torchmasked_fill_的「掩码处填充、其余保持」语义。从语义对照看torch 的x.masked_fill_(mask, value)等价于x torch.where(mask, value, x)而 PyPTO 骨架用gt先构造mask、再where三路选择二者完全对应。值得注意gt与where在 torch-pypto-op-mapping.md 中均属于「同名同参」的直接映射算子逐元素双输入分组无需额外适配即可使用。2.5 结果写回assemble 实现 inplacepypto.assemble(r, [i] [0] * len(inner), out)pypto.assemble将本次迭代的计算结果r写回输出张量out的对应位置偏移[i] [0] * len(inner)与前面的pypto.view切片位置一一对应。由于 inplace 版本的调用约定中out与a共享底层存储这一写回操作在效果上就是原地更新——这是 PyPTO 中实现 inplace 语义的标准手法。对比非 inplace 版本 masked_fill.md两者计算主体完全一致仅 Note 措辞不同inplace 版本额外注明「out 即 inplace 写回」。同样地mul_.md 也采用完全相同的模式pypto.mul计算结果直接assemble回out因为「pypto 无 inplace结果写回 out」。这说明「计算 assemble 写回 out」是 PyPTO 中表达所有 inplace 算子的通用骨架。3. 占位符约定与可运行性说明本文骨架是kernel 参考骨架reference skeleton不是可直接编译的生产模板。examples/README.md 明确说明每个op.md仅展示接口组合与轴切分模式哪些轴 loop、哪些轴整块loop 轴、unroll_list、tile shape、动态轴处理等需按实际 shape/dtype 与平台约束确定并调优骨架未逐一经 NPU 编译验证。骨架共用占位符占位符含义sl输入 shape 列表如[B, S, D]ol输出 shape 列表pypto_dtype元素 dtype如pypto.DT_FP32batch被 loop 的外层轴长度通常sl[0]inner单次迭代处理的内层 shape如sl[1:]将其替换为真实值即可得到最小可运行形态。例如对[B, D] [8, 128]的 FP32 输入batch 8、inner [128]、pypto_dtype pypto.DT_FP32。4. 动态 shape 与工程化注意事项将本骨架投入生产实现时需要对照 pypto-kernel-design-format.md 的分层规范与 SKILL.md 的硬约束清单注意以下几点动态 shape 兼容性gt、where这类逐元素 API 对动态维度的容忍度高于归约/矩阵乘类 API但 kernel 签名若使用动态轴仍需确认编译器能正确推断view切片与 tile 形状。若batch是动态的例如在pypto.loop中使用符号标量需确认运行时循环展开路径可行dtype 入口约束PyPTO 的 from_torch 入口支持 FP16/BF16/FP32/FP64 与多档整型及 BOOLmasked_fill_场景通常使用 FP32 或 BF16需确认所选 dtype 在gt/where/full上均受支持且fill_val如-1e9在目标 dtype 内可精确表示pypto.loop(1)的使用边界若把本骨架中 batch 循环折叠为单次迭代需注意pypto.loop(1)仅在 kernel 没有任何其他pypto.loop调用时才是合法的「layout-check 逃生口」存在真实 batch 循环时不应再套一层冗余的pypto.loop(1)对应 lint 规则 OL46tile shape 的作用域set_vec_tile_shapes影响其后所有 PyPTO 调用的 tile 布局。本骨架只有单一计算阶段在 loop 内设置一次即可若后续与其它 stage 融合建议把 tile 配置下沉到各pypto_*子 kernel 中对应 lint 规则 OL47。5. 真实场景印证Transformer 中的 masked_fill 用法masked_fill/masked_fill_组合方案在仓库的模型实现中有真实调用场景。以 src/pypto_gym/transformers/qwen3_5_9b/modeling_qwen3_5.py 为例其位置编码构造使用position_ids attention_mask.long().cumsum(-1) - 1 position_ids position_ids.masked_fill(attention_mask 0, 0)这里正是「先算掩码条件attention_mask 0再对掩码位置填充 0」的典型业务逻辑——与 kernel 骨架中「gt生成掩码 →where三路选择」的组合方案结构完全同构。类似的masked_fill调用还出现在 modeling_qwen3_5.py、modeling_gemma4.py、modeling_kimi.py 等众多模型的 attention mask 处理路径中。当这些模型算子需要下沉为 NPU kernel 时本文骨架即为masked_fill类逻辑的落地模板。6. 总结masked_fill_inplace的 PyPTO 参考骨架给出了一个简洁且可复用的逐元素算子范式轴切分batch 轴pypto.loop切分 pypto.view取切片last-dim 整块处理算子类型纯逐元素/比较 → Vector 型只需set_vec_tile_shapes计算组合gt生成掩码 full构造填充值 where三路选择等价于 torch 的masked_fill_inplace 语义通过把out与输入共享存储、assemble写回实现这是 PyPTO 表达所有 inplace 算子的通用方式。该骨架与 masked_fill.md非 inplace 版、mul_.md 等共同构成 pypto-api-explore 技能集中「torch 逐元素/索引算子 → PyPTO 组合方案」的参考族谱可直接作为算子开发与 API 探索阶段的起点。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考