PyTorch 1.13 GPU布尔索引出现numel整数乘法溢出问题求助
布尔掩码索引PyTorch GPU张量触发
numel: integer multiplication overflow错误的问题 问题背景
使用布尔掩码张量对PyTorch张量进行索引时,触发RuntimeError: numel: integer multiplication overflow错误,该错误在基于另一张量的布尔值过滤张量时出现。
- 环境:PyTorch 1.13 + CUDA 11.7
- 工作场景:基于MobileNetV2骨干的nanodet模型
调试发现,形状为(73, 3)的cost张量和形状为(73,)的布尔类型prior_match_gt_mask张量会引发该错误。构造的测试代码如下:
import torch a = torch.randn(73, 3) b = torch.randint(0, 2, (73,)).bool() print(a.cpu()[b.cpu(),:]) print(a.cuda()[b.cuda(), :])
CPU上的布尔索引操作无异常,但GPU上的布尔索引操作有时正常,有时触发上述错误。
问题解答
已知问题说明
PyTorch 1.13版本的GPU布尔索引确实存在一些已知的边缘场景bug,其中就包含这类小张量在特定掩码组合下的整数溢出问题,该问题在后续的PyTorch 2.0及以上版本中已经被修复。
规避方案
- 升级PyTorch版本
直接升级到PyTorch 2.0+版本,官方已经修复了该类GPU布尔索引的溢出问题,这是最彻底的解决方式。 - 替换为整数索引
如果暂时无法升级版本,可以将布尔掩码转换为整数索引来规避:
通过# 替代原GPU索引操作 mask_indices = torch.nonzero(b.cuda(), as_tuple=True)[0] result = a.cuda()[mask_indices, :]torch.nonzero获取掩码中True对应的位置索引,再用整数索引进行张量切片,避免布尔索引的底层实现问题。 - 调整张量维度(适用性较窄)
可以先将小张量进行维度扩展,比如添加一个空维度后再执行索引操作,但该方案仅适用于部分场景,可靠性不如前两种。
内容的提问来源于stack exchange,提问作者Walter Wang
相关产品推荐
相关产品推荐

