viable/strict/1789770346: [inductor] Optimize full slice_scatter chains (#195631)
PyTorch 发布 viable/strict/1789770346 版本说明,介绍 Inductor 的一项编译优化(PR #195631):把填满整个张量的 slice_scatter 链改写为峰值内存更高效的 cat 或 copy_。原文给出的两种输出为:当 Inductor 的 cat lowering 不会让输出与 chunk 分配重叠时用 aten.cat([chunk0, chunk1], dim);否则用 aten.copy_ 把各 chunk 写入复用的 base,使每个 chunk 在自己的写入处即释放。该变换要求完整覆盖以及兼容的 shape、stride、dtype、device、layout 等条件,存储同一性来自 FakeTensor 元数据而非算子白名单。
Development
- First Reportviable/strict/1789770346: [inductor] Optimize full slice_scatter chains (#195631)PyTorch Core
- Current Assessment判断:这类改动反映训练与推理框架的竞争重心正从算子覆盖转向编译期内存与调度效率,因为长上下文与分块(chunking)类负载的瓶颈常是峰值显存而非算力。对使用 PyTorch 编译栈的团队,这类 pass 可能在不改模型代码的前提下降低显存占用。可验证的下一信号:PyTorch 发布说明或 PR 中是否给出该 pass 的显存与吞吐对比,以及其他框架是否跟进同类 slice_scatter 链重写。Agent Pulse · analysis
PyTorch 在 viable/strict/1789770346 的发布说明中记录了一项 Inductor 编译优化(PR #195631):将填满整个张量的 slice_scatter 链改写为峰值内存更高效的 cat 或 copy_。背景是 functionalizing chunking 会产生一串 slice_scatter 节点,并在链尾持有全尺寸中间结果;该 pass 的目标是让 chunking 区域在峰值内存上最省。原文给出两种输出:一是 aten.cat([chunk0, chunk1], dim),前提是 Inductor 的 cat lowering 不会让输出与 chunk 分配重叠;二是把各 chunk 复制进复用的 base,使每个 chunk 在自己的写入处消亡。存储同一性判断来自 FakeTensor 元数据而非算子白名单,因此图输入视图不依赖构造它的视图算子;变换需要完整覆盖与兼容的 shape、stride、dtype、device、layout 等条件。
这是编译期图重写层面的内存优化,而非新算子或新硬件能力:把「先建空张量再逐段 slice_scatter」的链式写法,替换为 cat 或就地 copy_,从而避免在链尾同时持有全尺寸中间结果。值得注意的是存储同一性由 FakeTensor 元数据判定,而不是靠算子白名单,这让图输入视图的识别更稳健。可验证的下一信号:该 pass 是否在后续版本说明中扩展到非完整覆盖的 slice_scatter、是否补充峰值内存基准数据,以及 cat 与 copy_ 两条路径的触发条件是否被文档化。
判断:这类改动反映训练与推理框架的竞争重心正从算子覆盖转向编译期内存与调度效率,因为长上下文与分块(chunking)类负载的瓶颈常是峰值显存而非算力。对使用 PyTorch 编译栈的团队,这类 pass 可能在不改模型代码的前提下降低显存占用。可验证的下一信号:PyTorch 发布说明或 PR 中是否给出该 pass 的显存与吞吐对比,以及其他框架是否跟进同类 slice_scatter 链重写。
判断:对以显存为瓶颈的训练与推理服务,编译期消除全尺寸中间结果可能直接降低单卡可承载的分块规模上限,进而影响单位任务成本与可部署的上下文长度。可验证的下一信号:官方基准中该 pass 前后的峰值显存与吞吐数字,以及云厂商或框架发行版是否将其纳入默认编译配置。
判断:若该优化被证明稳定,chunking 类负载的显存峰值可能下降,使更大分块或更长序列在同等硬件上可行。可验证的下一信号:后续版本说明是否列出该 pass 的默认开启状态、覆盖条件放宽,以及是否出现回归报告。