PyTorch训练报错:'numpy.ndarray' object has no attribute 'to'
报错产生原因
.to(device)是PyTorch张量(Tensor)专属的设备迁移方法,numpy数组本身没有这个方法。
出现单张图像运行正常、2张及以上图像触发报错的核心原因是:train_generator数据生成器输出的batch数据类型不一致:单张输入时返回的是PyTorch Tensor类型,当batch包含多张图像时,数据加载管线没有执行numpy数组转Tensor的逻辑,直接返回了numpy.ndarray类型的features和targets,此时调用.to()方法就会抛出属性错误。
修复方案
两种修复方式按需选择:
- 快速修复:读取batch数据后,手动将numpy数组转换为PyTorch张量,再迁移到指定计算设备。替换训练循环中对应代码即可:
### 替换原有的features = features.to(device)、targets = targets.to(device)两行代码 import torch # 先将numpy数组转为Tensor,再迁移到目标设备 features = torch.from_numpy(features).to(device) targets = torch.from_numpy(targets).to(device) # 注意:如果图像numpy数组为通道在后格式(shape为[batch数, 图像高, 图像宽, 通道数]), # 需要额外调整维度顺序适配PyTorch模型要求的通道在前格式,取消下一行注释即可 # features = features.permute(0, 3, 1, 2)
- 根源修复:检查数据加载管线配置,确认数据预处理transform中加入了
torchvision.transforms.ToTensor()转换逻辑,保证数据生成器不管输出单张还是多张图像的batch,返回的都是标准PyTorch Tensor类型,后续不需要手动做类型转换。
内容的提问来源于stack exchange,提问作者Anak Cerdas
相关产品推荐
相关产品推荐

