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

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} ]

错误原因分析

  1. Dataset设计逻辑错误:PyTorch的Dataset应返回单个样本,而非手动划分批量,DataLoader会自动完成批量组装。原MyDataGenerator手动按batch_size切分数据,导致样本和通道维度混淆。
  2. 张量维度处理错误:原输入inputs维度为(1, 181, 512, 512)(1为单样本批量,181为通道数),使用torch.squeeze(data)后丢失批量维度,后续索引将通道维度误当作样本数量。
  3. Loss函数未实现:原hs_loss仅返回模型输出,未计算损失值,导致训练逻辑失效。
  4. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 02:40:58