PyTorch1.9调用normalize报IndexError维度超出范围错误如何解决
错误根因定位
本次报错的核心原因是调用torch.nn.functional.normalize时传入的dim参数和张量实际维度不匹配:
- 从报错信息
Dimension out of range (expected to be in range of [-2, 1], but got 2)可以直接判断:你的masked_embedding是2维张量,合法的维度取值范围只有-2、-1、0、1,不存在维度2 - 你原来的旧版L2归一化代码中,计算norm用的是
dim=1,说明业务逻辑需要在第1个维度(特征维度)上做归一化,修改代码时误把dim改成了2,直接导致维度越界
修复方案
直接把dim参数改成和旧版代码一致的dim=1即可,修改后的代码如下:
if self.l2_norm: masked_embedding = torch.nn.functional.normalize(masked_embedding, p=2.0, dim=1, eps=1e-10, out=None)
你原来旧版代码在高版本PyTorch报错的原因是旧版写法没有加keepdim=True,高版本PyTorch对维度不匹配的张量除法的广播规则更严格,用官方normalize接口已经默认处理了维度对齐的问题,只要dim参数和旧版保持一致就可以完全兼容原有逻辑。
可选优化(消除告警)
运行日志中的transforms.Scale弃用警告,只需要全局搜索代码里的transforms.Scale替换成transforms.Resize即可,二者功能完全一致。
内容的提问来源于stack exchange,提问作者Mona Jalal
相关产品推荐
相关产品推荐

