PyTorch调用index_add函数后出现permute维度不匹配RuntimeError求助
问题原因
这个问题属于PyTorch特定版本的框架bug,和你的代码逻辑无关,你传入的参数完全符合index_add的接口规范。
该bug出现在启用CUDA侧确定性算法的场景下,PyTorch 1.12.x~2.0.x的部分小版本中,index_add的CUDA确定性实现内部存在逻辑错误:处理dim=0的加法时,内部代码错误地对2维张量执行了不符合维度数的permute操作,才会抛出number of dims don't match in permute的报错,和你实际传入的参数没有关系。
验证与解决方案
- 验证方法:临时注释掉
torch.use_deterministic_algorithms(True)这行配置,再次运行相同代码,如果可以正常执行,即可确认是该确定性实现的bug。 - 解决方案
- 优先升级PyTorch到2.1.0及以上版本,该bug已经在后续版本的官方补丁中被修复。
- 如果暂时无法升级版本,可以临时将运算转移到CPU执行,完成后再迁回CUDA,示例代码如下:
# 临时迁移到CPU执行规避bug Y = Y.index_add(0, indices.cpu(), X.cpu()).cuda()
- 如果必须保留CUDA侧执行和确定性配置,可以用功能等价的`scatter_add`接口替代:
Y = Y.scatter_add(0, indices.unsqueeze(-1).expand_as(X), X)
内容的提问来源于stack exchange,提问作者thebesttony
相关产品推荐
相关产品推荐

