ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

PyPTO masked_fill 算子 kernel 参考骨架:基于 gt + where 组合的掩码填充实现

PyPTO masked_fill 算子 kernel 参考骨架:基于 gt + where 组合的掩码填充实现 PyPTO masked_fill 算子 kernel 参考骨架基于 gt where 组合的掩码填充实现【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gymmasked_fill 是深度学习模型中高频出现的逐元素掩码填充算子典型场景如 attention 的 causal mask 填充在 PyPTO 编程框架下没有同名原生 API而是通过gt生成掩码、where条件选择两个 Vector 级原语组合实现。本文基于 PyPTO-Gym 仓库中pypto-api-exploreskill 提供的 kernel 参考骨架完整讲解 masked_fill 的 PyPTO 实现方式、batch 轴 loop 切分与 last-dim 整块策略、占位符替换规则并对比 inplace 变体与 masked_scatter 的同类组合方案帮助算子在 NPU 上以 Vector 流水线快速落地。1. 算子语义与 PyPTO 组合方案masked_fill(input, mask, value)的语义是当mask对应位置为 True 时将input中该位置替换为value否则保留原值。在 PyTorch 中它常与布尔掩码配合用于 attention mask、padding mask 等场景。PyPTO 没有提供名为masked_fill的直接 API但该算子的逐元素语义可以精确分解为两个原子操作pypto.gt(a_s, zero)比较生成布尔掩码pypto.where(mask, fill_val, a_s)按掩码做条件选择即掩码为 True 取填充值否则取原值。这一分解在 torch-pypto-op-mapping.md 的命名映射表中被明确记录Torch 算子PyPTO 组合方案kernel 参考骨架masked_fillgtwheremasked_fill.mdmasked_fill_inplacegtwheremasked_fill_inplace.mdmasked_scatterwherecumsumcastgathermasked_scatter.md可见where是这类掩码类算子的公共底座masked_fill 用gt现造掩码masked_scatter 则直接消费外部传入的mask张量其余流程几乎一致见下文 4.2 节对比。2. 完整 kernel 参考骨架以下代码节选自 examples/masked_fill.md为 masked_fill 的标准参考骨架pypto.frontend.jit(runtime_options{run_mode: pypto.RunMode.NPU}) def masked_fill_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)骨架顶部注释点明了该算子的核心切分策略batch 轴 loop 切分last-dim 整块gt 生成掩码后 where 填充。逐行拆解其执行流程JIT 入口pypto.frontend.jit(runtime_options{run_mode: pypto.RunMode.NPU})声明以 NPU 运行模式编译该 kernelbatch 轴切分pypto.loop(batch, namebatch, unroll_list[1])将 batch 维作为外层循环逐行处理unroll_list[1]表示该轴不展开、每次迭代只处理一行切片视图pypto.view(a, [1] inner, [i] [0] * len(inner))以零拷贝视图方式取出第i个 batch 行shape 变为[1] inner起始偏移为[i] [0]*len(inner)Vector Tile 声明pypto.set_vec_tile_shapes(1, *inner)告知编译器按1 × inner的 tile 形状执行 Vector 计算构造掩码zero pypto.full([1] inner, 0.0, pypto_dtype)生成全 0 常量张量mask pypto.gt(a_s, zero)逐元素比较得到布尔掩码元素 0 为 True构造填充值fill_val pypto.full([1] inner, -1e9, pypto_dtype)生成全-1e9的填充张量条件选择r pypto.where(mask, fill_val, a_s)完成核心填充逻辑——掩码为 True 的位置取-1e9否则保留原值写回pypto.assemble(r, [i] [0] * len(inner), out)将本行结果回写到输出的对应位置。3. 占位符替换把骨架变成可运行 kernel参考骨架中使用了通用占位符替换规则在 examples/README.md 中统一约定占位符含义示例取值sl输入 shape 列表[B, S, D]ol输出 shape 列表[B, S, D]pypto_dtype元素 dtypepypto.DT_FP32batch被 loop 的外层轴长度通常sl[0]Binner单次迭代处理的内层 shapesl[1:]结合 README 给出的最小可运行 setup 即可组装出完整 kernelimport pypto B, D 8, 128 sl, ol [B, D], [B, D] pypto_dtype pypto.DT_FP32 batch, inner B, [D]替换后得到完整的 8×128 输入示例外层循环迭代 8 个 batch 行每行以[1, 128]的 Vector tile 整块计算。注意骨架中的-1e9填充值、0.0阈值和unroll_list均为示意需按实际业务语义如 attention mask 中的-inf近似与平台约束调整。4. 同类变体与对比4.1 inplace 版本masked_fill_examples/masked_fill_inplace.md 提供了 inplace 语义的骨架代码与普通版本完全同构唯一的语义差异体现在注释中batch 轴 loop 切分last-dim 整块gt 生成掩码后 where 填充out 即 inplace 写回。即在调用侧将a与out绑定为同一张量由assemble完成原地写回。PyPTO 的viewassemble数据流天然支持这一用法无需引入额外临时拷贝。4.2 masked_scatterwhere 组合的另一种形态examples/masked_scatter.md 展示了where的兄弟用法掩码由外部传入而非gt现造for i in pypto.loop(batch, namerow, unroll_list[1]): x_s pypto.view(x, [1] inner, [i] [0] * len(inner)) mask_s pypto.view(mask, [1] inner, [i] [0] * len(inner)) src_s pypto.view(source, [1] inner, [i] [0] * len(inner)) pypto.set_vec_tile_shapes(1, *inner) r pypto.where(mask_s, src_s, x_s) pypto.assemble(r, [i] [0] * len(inner), out)对比可见两者共享相同的 batch loop、tile shape 与where三参数选择结构差异仅在掩码来源gt计算 vs 外部张量与第二操作数常量fill_valvs 数据张量src_s。这印证了gtwhere是 PyPTO 掩码类算子的通用范式。5. 仓库中的真实使用场景与约束5.1 attention mask 填充的真实调用masked_fill 在仓库中的典型落地场景是因果注意力掩码构造。在 modeling_qwen3.py 中可以看到 PyTorch 侧的等价语义causal_mask[:, :, :, :mask_length] causal_mask[:, :, :, :mask_length].masked_fill(...)即在 causal mask 的指定切片上用掩码填充通常填充为极小值这与骨架中fill_val -1e9的设定完全对应——attention 场景常用-inf或足够小的负数作为被掩码位置的填充值使 softmax 后趋近于 0。qwen3_5、gemma4、kimi、spatial_ssrl、gutenocr、phi_3 等多个模型的 modeling 文件中也存在同类调用说明该算子模式在仓库覆盖的模型族中具有普遍性。5.2 设计约束骨架不是标准模板需要特别强调的是参考骨架仅展示接口组合与轴切分模式并不代表已通过 NPU 编译验证的生产实现。以下几点必须在实际开发中重新确认loop 轴、unroll_list、tile shape 需按实际 shape / dtype 与平台约束调优examples/README.md 明确说明骨架未逐一经 NPU 编译验证skew 检查与 lint 门禁优先pypto-api-exploreskill 定位该目录为API 用法参考而非 production 标准当写法与 lint / 门禁冲突时以 lint 为准上游禁止项在 pypto-kernel-design-format.md 中masked_fill本身被列为如果使 NPU lowering 复杂化、与 PyPTO 路径不一致则应禁止的算子——这正是它必须被显式分解为gtwhere组合、而不是作为高层算子直译的原因动态 shape 注意计算类 API含gt、where在编译期需要 concrete shape若batch或inner含动态维需采用 loop 切 tile 策略并在风险评估中标注。5.3 与 pypto-api-explore 工作流的衔接该骨架在pypto-api-exploreskill 的 API 探索工作流中承担本地映射优先的入口作用当收到masked_fill类算子需求时先在 torch-pypto-op-mapping.md 命中gtwhere组合条目再读取本骨架确认接口组合与切分模式之后仍需通过 Explore subagent 核实具体约束与生产实现。因此本文骨架既是可复用的 kernel 起点也是 API 映射结论的直接证据。6. 总结masked_fill 在 PyPTO 中的最佳实践可以概括为一条清晰路径语义分解masked_fill拆解为gt造掩码where条件填充两个 Vector 原语切分策略batch 轴pypto.loop外层切分last-dim 以set_vec_tile_shapes整块计算view/assemble完成切片与回写参数落实按实际 shape / dtype 替换占位符fill_val依业务语义取值attention 场景常用-1e9级极小值并对 loop 轴、unroll 与 tile shape 做平台级调优变体复用inplace 版本仅需将out与a绑定masked_scatter 则复用同一where底座、替换掩码与数据来源。参考该骨架开发者可以在 PyPTO-Gym 的pypto-api-explore框架下快速生成可编译、可调优的 NPU masked_fill kernel并以此类推到其他掩码类算子。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表