PyTorch设备不匹配报错排查:DenseNet模型GPU训练异常
问题分析:GPU部署DenseNet时的设备不匹配错误
使用kuangliu的GitHub仓库中DenseNet训练CIFAR-10的代码,部署到GPU加速训练时触发RuntimeError: Expected all tensors to be on the same device, but found at least two devices, cpu and cuda:0错误。
错误原因
问题出在test()函数中:模型已经通过.to(device)部署到GPU,但输入张量x仍生成在CPU上,导致模型计算时,GPU上的模型参数与CPU上的输入张量无法匹配运算。
解决方法
将输入张量x同步到与模型相同的设备上,修改test()函数中的输入生成代码:
修改前:
x = torch.randn(1,3,32,32)
修改后:
x = torch.randn(1,3,32,32).to(device)
完整修改后的test()函数:
def test(): device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') net = densenet_cifar().to(device) x = torch.randn(1,3,32,32).to(device) y = net(x) print(y)
额外说明
你提供的自定义Bottleneck、Transition、DenseNet类本身没有问题,所有模型参数在调用.to(device)后都会正确迁移到GPU,错误仅由输入张量未同步设备导致。
内容的提问来源于stack exchange,提问作者sheep-coder
相关产品推荐
相关产品推荐

