PyTorch批量训练报错:通道不匹配与张量维度异常求助
批量训练通道不匹配与维度异常问题解决
问题描述
单样本模式下PyTorch代码可正常运行,但尝试批量训练时:
- 设置
batch_size=4,报错:Given groups=1, weight of size [16, 9, 1, 1], expected input[1, 4, 512, 512] to have 9 channels, but got 4 channels instead - 设置
batch_size=9,报错:3D or 4D (batch mode) tensor expected for input, but got: [ torch.cuda.FloatTensor{1,0,512,512} ]
错误原因分析
- Dataset设计逻辑错误:PyTorch的
Dataset应返回单个样本,而非手动划分批量,DataLoader会自动完成批量组装。原MyDataGenerator手动按batch_size切分数据,导致样本和通道维度混淆。 - 张量维度处理错误:原输入
inputs维度为(1, 181, 512, 512)(1为单样本批量,181为通道数),使用torch.squeeze(data)后丢失批量维度,后续索引将通道维度误当作样本数量。 - Loss函数未实现:原
hs_loss仅返回模型输出,未计算损失值,导致训练逻辑失效。 - Targets未适配批量:原
targets为单样本数据,批量训练时需对应生成批量标签。
修正后的完整代码
import torch import torch.nn as nn import numpy as np device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') # attention module for FRU class attention_FRU(nn.Module): def __init__(self, num_channels_down, pad='reflect'): super(attention_FRU, self).__init__() # layers to generate conditional convolution weights self.gen_se_weights1 = nn.Sequential( nn.Conv2d(num_channels_down, num_channels_down, 1, padding_mode=pad), nn.LeakyReLU(0.2, inplace=True), # Dont use Softplus here nn.Sigmoid()) # create conv layers self.conv_1 = nn.Conv2d(num_channels_down, num_channels_down, 1, padding_mode=pad) self.norm_1 = nn.BatchNorm2d(num_channels_down, affine=False) self.actvn = nn.LeakyReLU(0.2, inplace=True) # self.actvn = nn.Softplus() def forward(self, guide, x): se_weights1 = self.gen_se_weights1(guide) dx = self.conv_1(x) dx = self.norm_1(dx) dx = torch.mul(dx, se_weights1) out = self.actvn(dx) return out class hs_net(nn.Module): def __init__(self, ym_channel, yh_channel, num_channels_down, num_channels_up, num_channels_skip, filter_size_down, filter_size_up, filter_skip_size): super(hs_net,self).__init__() self.FRU = attention_FRU(num_channels_down) self.up_bic = nn.Upsample(scale_factor=4, mode='bicubic') self.up_trans = nn.ConvTranspose2d(yh_channel,yh_channel,filter_size_down,stride=4,padding=1) self.guide_ms = nn.Sequential( nn.Conv2d(ym_channel, num_channels_down, filter_size_down, padding ='same',padding_mode='reflect'), nn.BatchNorm2d(num_channels_down), # nn.LeakyReLU(0.2)) nn.Softplus()) self.enc = nn.Sequential( nn.Conv2d(num_channels_down, num_channels_down, filter_size_down,padding='same', padding_mode='reflect'), nn.BatchNorm2d(num_channels_down), # nn.LeakyReLU(0.2)) nn.Softplus()) self.skip = nn.Sequential( nn.Conv2d(num_channels_down, num_channels_skip, filter_skip_size, padding ='same', padding_mode='reflect'), nn.BatchNorm2d(num_channels_skip), # nn.LeakyReLU(0.2)) nn.Softplus()) self.dc = nn.Sequential( nn.Conv2d((num_channels_skip + num_channels_up), num_channels_up, filter_size_up,padding='same',padding_mode='reflect'), nn.BatchNorm2d(num_channels_up), # nn.LeakyReLU(0.2)) nn.Softplus()) self.out_layer = nn.Sequential( nn.Conv2d(num_channels_up, yh_channel, 1, padding_mode='reflect'), nn.Sigmoid()) self.conv_hs = nn.Sequential( nn.Conv2d(yh_channel,num_channels_down,filter_size_down, padding = 'same',padding_mode = 'reflect')) # nn.BatchNorm2d(num_channels_down), # nn.Softplus()) self.conv_bn = nn.Sequential( nn.Conv2d(num_channels_down,num_channels_down,filter_size_down, padding = 'same',padding_mode = 'reflect'), # nn.BatchNorm2d(num_channels_down), # # nn.LeakyReLU(0.2)) nn.Softplus()) self.ym_channels= ym_channel def forward(self, inputs): ym = inputs[:, :self.ym_channels, :, :] yh = inputs[:, self.ym_channels:, :, :] ym_en0 = self.guide_ms(ym) ym_en1 = self.enc(ym_en0) ym_en2 = self.enc(ym_en1) ym_en3 = self.enc(ym_en2) ym_en4 = self.enc(ym_en3) ym_dc0 = self.enc(ym_en4) ym_dc1 = self.enc(ym_dc0) ym_dc2 = self.dc(torch.cat((self.skip(ym_en4), ym_dc1), dim=1)) ym_dc3 = self.dc(torch.cat((self.skip(ym_en3), ym_dc2), dim=1)) ym_dc4 = self.dc(torch.cat((self.skip(ym_en2), ym_dc3), dim=1)) ym_dc5 = self.dc(torch.cat((self.skip(ym_en1), ym_dc4), dim=1)) ym_dc6 = self.dc(torch.cat((self.skip(ym_en0), ym_dc5), dim=1)) yh_6 = self.FRU(self.conv_hs(yh), ym_dc0) yh_7 = self.FRU(self.conv_bn(yh_6), ym_dc1) yh_8 = self.FRU(self.conv_bn(yh_7), ym_dc2) yh_9 = self.FRU(self.conv_bn(yh_8), ym_dc3) yh_10 = self.FRU(self.conv_bn(yh_9), ym_dc4) yh_11 = self.FRU(self.conv_bn(yh_10), ym_dc5) yh_12 = self.FRU(self.conv_bn(yh_11), ym_dc6) out = self.out_layer(yh_12) # 上采样到目标hsi的尺寸(128x128) out = self.up_bic(out) return out # 修正后的Dataset:返回单个样本,由DataLoader自动组装批量 class MyDataGenerator(torch.utils.data.Dataset): def __init__(self, data): self.data = data # 数据维度:(num_samples, 181, 512, 512) def __len__(self): return self.data.shape[0] def __getitem__(self, index): return self.data[index] # 返回单个样本:(181, 512, 512) # 生成批量输入:4个样本,每个样本通道数181(9+172) batch_size = 4 inputs = torch.from_numpy(np.random.rand(batch_size, 181, 512, 512)).to(device,dtype=torch.float) num_iter = 1000 LR = 0.001 n_channels=172 net=hs_net(ym_channel=9, yh_channel=n_channels, num_channels_down=16, num_channels_up=16, num_channels_skip=16, filter_size_down=1, filter_size_up=1, filter_skip_size=1).to(device) # 生成批量targets:对应每个输入样本 msi_batch = torch.from_numpy(np.random.rand(batch_size, 9, 512, 512)).to(device,dtype=torch.float) hsi_batch = torch.from_numpy(np.random.rand(batch_size, 172, 128, 128)).to(device,dtype=torch.float) targets = [msi_batch, hsi_batch] optimizer = torch.optim.Adam(net.parameters(), lr=LR, eps=1e-3, amsgrad=True) loss_fn = nn.MSELoss() # 定义损失函数,可根据任务调整 def hs_loss(model, inputs, targets): yh_target = targets[1] xhat = model(inputs) # 计算损失:模型输出与目标hsi的MSE loss = loss_fn(xhat, yh_target) return loss, xhat # 批量训练逻辑 train_set = MyDataGenerator(inputs) data_generator = torch.utils.data.DataLoader(train_set, batch_size=batch_size, shuffle=True) for it in range(num_iter): optimizer.zero_grad() total_loss = 0.0 for batch in data_generator: # batch维度:(batch_size, 181, 512, 512) loss, out_HR = hs_loss(net, batch, targets) total_loss += loss.item() loss.backward() optimizer.step() if (it+1) % 50 == 0: print(f"Iteration {it+1}, Average Loss: {total_loss/len(data_generator):.6f}")
关键改动说明
- 重构Dataset:删除
batch_size参数,__getitem__返回单个样本,__len__返回总样本数,由DataLoader负责批量组装。 - 调整输入/Target维度:生成批量输入
(N,181,512,512)和对应批量Target(N,9,512,512)、(N,172,128,128),匹配模型输入要求。 - 修复Loss函数:添加MSE损失计算(可根据任务替换为其他损失),返回损失值和模型输出,满足训练逻辑。
- 补充上采样步骤:模型输出需上采样到目标HSI的尺寸(128x128),否则损失计算时维度不匹配。
内容的提问来源于stack exchange,提问作者George
相关产品推荐
相关产品推荐

