You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

训练图像重建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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.23 22:42:32