训练图像重建CNN遇RuntimeError:Input与bias类型不匹配(double/float)
问题原因
这个错误是因为输入数据的类型(double/float64)和模型参数的类型(float32)不匹配导致的。PyTorch要求张量运算时,参与计算的所有张量(包括模型的权重、偏置)必须是相同的数据类型。
你的代码里:
- 手动将
sys_mat转成了float64,后续生成的投影数据true_sinogram_torch也继承了这个类型 - 而PyTorch的卷积层、激活层等默认使用
float32类型的参数,当输入float64的数据到模型时,就会触发类型不匹配的报错
修复步骤
1. 统一系统矩阵的数据类型为float32
找到代码中设置sys_mat dtype的行,把float64改成float32:
# 原代码 sys_mat = sys_mat.to(torch.float64) # 修改为 sys_mat = sys_mat.to(torch.float32)
2. 确保输入张量为float32
修改np_to_00torch函数,让生成的张量默认是float32类型:
# 原代码 def np_to_00torch(np_array): return torch.from_numpy(np_array).unsqueeze(0).unsqueeze(0) # 修改为 def np_to_00torch(np_array): return torch.from_numpy(np_array).unsqueeze(0).unsqueeze(0).to(torch.float32)
3. 修正CNN forward函数的缩进错误
原代码中return x前面有多余的缩进,会导致语法错误,修正后:
def forward(self,x): x = torch.squeeze(self.CNN(x.unsqueeze(0).unsqueeze(0))) return x
可选:修复设备设置的笔误
代码中设备设置行有个笔误:cuda:o应该是cuda:0,虽然不影响CPU运行,但用GPU时会出错:
# 原代码 device = torch.device("cuda:o" if torch.cuda.is_available() else "cpu"); print(device) # 修改为 device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu"); print(device)
关键修改后的代码片段
# 修正设备设置 device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu"); print(device) # 修正np_to_00torch函数 def np_to_00torch(np_array): return torch.from_numpy(np_array).unsqueeze(0).unsqueeze(0).to(torch.float32) # 修正sys_mat dtype sys_mat = make_torch_system_matrix(nxd, nrd, nphi).to(device) sys_mat = sys_mat.to(torch.float32) # 修正CNN forward函数 class CNN(nn.Module): def __init__(self, num_channels): super(CNN, self).__init__() self.CNN = nn.Sequential( nn.Conv2d(1, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, num_channels, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), nn.Conv2d(num_channels, 1, 3, padding=(1, 1), padding_mode='reflect'), nn.PReLU(), ) def forward(self,x): x = torch.squeeze(self.CNN(x.unsqueeze(0).unsqueeze(0))) return x
内容的提问来源于stack exchange,提问作者Dominic Cugliari
相关产品推荐
相关产品推荐

