已合并
fix: 修复 clamp_npu_output_size 空张量短路导致 broadcast shape 计算错误 #5687
wangqi_ai创建于 14 天前
fix: 修复 clamp_npu_output_size 空张量短路导致 broadcast shape 计算错误 #5687
已合并
Pull Request已成功合入, 合并人@ascend-robot
(感谢 wangqi_ai 的贡献)ascend-robot
14 天前 评论:
14 天前 评论:
atomgit-bot
14 天前 评论:
14 天前 评论:
变更摘要
此 PR 修复了 clamp_npu_output_size 函数中因空张量短路分支导致的 broadcast shape 计算错误。当 self.numel() == 0 时,原实现直接返回 self.sizes(),跳过了与 min/max 的 broadcast 维度推导,导致维度扩展缺失和广播合法性校验被吞掉。修复方案是删除该短路分支(7 行代码),使空张量场景统一经 broadcast_ops_npu_output_size(底层调用 at::infer_size)计算输出 shape,与 PyTorch 上游 torch.clamp.Tensor 的 broadcast 语义保持一致。
主要改动
- 删除
self.numel() == 0短路返回逻辑:移除KernelNpuOutputSize.cpp中clamp_npu_output_size函数内对self.numel() == 0的特判分支,该分支原本直接按self.dim()逐维拷贝尺寸并返回,不再保留。 - 空张量输出 shape 计算统一走 broadcast 路径:删除短路后,即使
self为空张量,也会进入后续的broadcast_ops_npu_output_size调用链进行正确的维度广播和合法性校验,消除维度数错误(如(0,)vs(3,0))和非法输入静默通过的问题。


atomgit-bot
14 天前 评论:
14 天前 评论:
14 天前 添加了label:ascend-cla/yes
此处折叠了41条消息 查看更多
梁松伟
11 天前 评论:
11 天前 评论:
/approve


11 天前 添加了label:approvedlgtm
11 天前 删除了label:ci-pipeline-passed
11 天前 关闭了关联的issue
11 天前 合入了pull request
【合入来源】
【修改方案】
问题根因
clamp_npu_output_size(op_plugin/utils/KernelNpuOutputSize.cpp)在self.numel() == 0时直接返回self.sizes(),跳过了与 min/max 的 broadcast 计算。该函数通过
op_plugin_functions.yaml中size: clamp_npu_output_size(self, min, max)配置被 gen_opapi 自动生成的clamp.Tensor/clamp.Tensor_out调用,bug 会真实触发。broadcast_ops_npu_output_size内部使用at::infer_size,已正确处理 0 维、left-pad 和 broadcast 合法性校验。短路分支绕过了这些能力,导致两类问题:shape 计算错误:当
self.ndim < min/max.ndim(broadcast 应扩展维度)或同维度但某维需扩展时,返回的 shape 维度数/数值错误。self=(0,),min=(3,1)→ 正确结果(3, 0),短路返回(0,)self=(1, 0),min=(3, 1)→ 正确结果(3, 0),短路返回(1, 0)吞掉 broadcast 合法性校验:不可 broadcast 的非法输入(如
self=(0,),min=(3,))被静默接受,错误延迟到后续 kernel 执行。修复方法
删除
self.numel() == 0短路分支(7 行),统一走broadcast_ops_npu_output_size(at::infer_size),与 PyTorch 上游torch.clamp.Tensor的 broadcast 语义保持一致。【资料变更】
不涉及
【接口变更】
不涉及
【功能验证】
验证场景(self.numel()==0 且需 broadcast):
(0,)(3, 1)(0,)❌(3, 0)✅(3, 0)(1, 0)(3, 1)(1, 0)❌(3, 0)✅(3, 0)(0,)(3,)(0,)❌(吞错)(0,)()标量(0,)✅(0,)✅(0,)【CheckList】
【本地验证】
修改后,torch.clamp和torch._refs.clamp_max均执行成功:
