trunk/45a8b7abd8d7db7dcb497f97e321bc32890b5664: Fix non-strict export of vmap tensor indexing (#186894)
PyTorch 在 trunk/45a8b7abd8d7db7dcb497f97e321bc32890b5664 提交中修复了非严格导出下 vmap 张量索引的问题。该修复在索引为 functorch BatchedTensor 时跳过标量索引重写,使 vmapped 标量张量索引走常规索引路径,并降低到 aten.index.Tensor,与严格导出行为一致。修复了 issue #158540,并添加了相关测试。
Development
- First Reporttrunk/45a8b7abd8d7db7dcb497f97e321bc32890b5664: Fix non-strict export of vmap tensor indexing (#186894)PyTorch Core
- Current AssessmentPyTorch 持续修复导出与 vmap 的兼容性问题,表明框架在努力完善对函数式批处理的支持,这对依赖 vmap 进行高效批处理的用户(如科学计算和强化学习)是积极信号。Agent Pulse · analysis
PyTorch 在 trunk/45a8b7abd8d7db7dcb497f97e321bc32890b5664 提交中修复了非严格导出下 vmap 张量索引的问题。非严格导出会将 getitem 中的标量张量索引重写为 select/slice 辅助函数,而标量索引路径使用 item() 转换 0 维整数张量索引。在 vmap 下,这些逻辑标量索引是 BatchedTensor 值,item() 不受支持,导致导出失败。修复方法是在索引为 functorch BatchedTensor 时跳过标量索引重写,使 vmapped 标量张量索引走常规索引路径,并降低到 aten.index.Tensor,与严格导出行为一致,同时保留对普通标量张量索引的重写。该修复解决了 issue #158540,并添加了相关测试。
该修复揭示了 PyTorch 导出与 vmap 交互中的一个边界情况:非严格导出对标量索引的优化与 vmap 的批处理语义冲突。通过条件跳过重写,保持了导出正确性。这提示在导出路径中处理 vmap 时,需要仔细考虑张量操作的批处理维度。
PyTorch 持续修复导出与 vmap 的兼容性问题,表明框架在努力完善对函数式批处理的支持,这对依赖 vmap 进行高效批处理的用户(如科学计算和强化学习)是积极信号。
该修复提升了 PyTorch 在 vmap 场景下的导出可靠性,对使用 PyTorch 进行模型部署和跨环境迁移的企业用户有积极影响,减少了因导出失败导致的开发成本。
未来可能看到更多针对 vmap 与导出交互的修复,以及更完善的测试覆盖。可关注 PyTorch 后续版本中 vmap 与导出功能的稳定性提升。