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

自定义数据集训练PyTorch CNN遇输入与偏置类型错误求助

PyTorch CNN训练类型不匹配问题排查

问题描述

参照微软PyTorch CNN教程搭建模型,未使用CIFAR-10数据集,改用自定义ASLDataset训练时,出现输入类型与偏置类型不匹配的错误,排查资料后未找到有效解决方法,请求协助定位问题。

相关代码

自定义数据集类

class ASLDataset(torch.utils.data.Dataset):
    def __init__(self, csv_file, root_dir="", transform=None):
        self.annotation_df = pd.read_csv(csv_file)
        self.root_dir = root_dir
        self.transform = transform

    def __len__(self):
        return len(self.annotation_df)

    def __getitem__(self, idx):
        image_path = os.path.join(self.root_dir, self.annotation_df.iloc[idx, 1])
        image = cv2.imread(image_path)
        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
        class_index = self.annotation_df.iloc[idx, 3]
        if self.transform:
            image = self.transform(image)
        return image, class_index

train_dataset = ASLDataset('./train.csv')
train_dataloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=num_workers)

val_dataset = ASLDataset('./test.csv')
val_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=num_workers)

classes = ('A', 'B', 'C', 'D', 'E', 'F', 'G', 'H', 'I', 'J', 'K', 'L', 'M', 'N', 'nothing', 'O', 'P', 'Q', 'R', 'S', 'space', 'T', 'U', 'V', 'W', 'X', 'Y', 'Z')

网络结构代码

class Network(nn.Module):
    def __init__(self):
        super(Network, self).__init__()

        self.conv1 = nn.Conv2d(in_channels=3, out_channels=12, kernel_size=5, stride=1, padding=1)
        self.bn1 = nn.BatchNorm2d(12)
        self.conv2 = nn.Conv2d(in_channels=12, out_channels=12, kernel_size=5, stride=1, padding=1)
        self.bn2 = nn.BatchNorm2d(12)
        self.pool = nn.MaxPool2d(2, 2)
        self.conv4 = nn.Conv2d(in_channels=12, out_channels=24, kernel_size=5, stride=1, padding=1)
        self.bn4 = nn.BatchNorm2d(24)
        self.conv5 = nn.Conv2d(in_channels=24, out_channels=24, kernel_size=5, stride=1, padding=1)
        self.bn5 = nn.BatchNorm2d(24)
        self.fc1 = nn.Linear(24 * 10 * 10, 10)

    def forward(self, input):
        output = F.relu(self.bn1(self.conv1(input)))
        output = F.relu(self.bn2(self.conv2(output)))
        output = self.pool(output)
        output = F.relu(self.bn4(self.conv4(output)))
        output = F.relu(self.bn5(self.conv5(output)))
        output = output.view(-1, 24 * 10 * 10)
        output = self.fc1(output)

        return output

训练代码片段

def train(num_epochs):
    best_accuracy = 0.0

    device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
    print("The model will be running on", device, "device")
    model.to(device)

    for epoch in range(num_epochs):
        running_loss = 0.0
        running_acc = 0.0

        for i, (images, labels) in enumerate(train_dataloader, 0):
            images = Variable(images.to(device))
            print(type(labels))
            labels = Variable(labels.to(device))

            optimizer.zero_grad()
            outputs = model(images)
            loss = loss_fn(outputs, labels)
            loss.backward()
            optimizer.step()

# 后续代码省略

错误信息

RuntimeError: 输入类型(torch.cuda.FloatTensor)与偏置类型(torch.FloatTensor)必须一致
(或反向:输入在CPU,偏置在GPU)

解决方案

1. 修复数据集输出的Tensor格式与类型

cv2读取的图像是numpy数组(HWC格式、uint8类型),PyTorch卷积层要求输入为CHW格式的float32 Tensor,且标签需为long类型用于分类任务。修改__getitem__方法:

def __getitem__(self, idx):
    image_path = os.path.join(self.root_dir, self.annotation_df.iloc[idx, 1])
    image = cv2.imread(image_path)
    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
    # 转换为Tensor,调整维度为CHW,归一化并转float32
    image = torch.from_numpy(image).permute(2, 0, 1).float() / 255.0
    if self.transform:
        image = self.transform(image)
    # 确保标签为long类型
    class_index = torch.tensor(self.annotation_df.iloc[idx, 3], dtype=torch.long)
    return image, class_index

2. 移除冗余的Variable并确保设备与类型对齐

PyTorch 0.4+版本后Variable已被废弃,直接将数据移到对应设备并指定类型:

# 替换训练代码中的数据转移部分
images = images.to(device, dtype=torch.float32)
labels = labels.to(device, dtype=torch.long)

3. 修正全连接层输入维度匹配问题

原网络中fc1的输入维度24*10*10是假设特征图尺寸为10x10,但实际需根据输入图像尺寸计算。若不确定输入尺寸,可改用自适应池化固定输出尺寸:

# 在Network类的__init__中添加自适应池化层
self.adaptive_pool = nn.AdaptiveAvgPool2d((10, 10))

# 修改forward方法
def forward(self, input):
    output = F.relu(self.bn1(self.conv1(input)))
    output = F.relu(self.bn2(self.conv2(output)))
    output = self.pool(output)
    output = F.relu(self.bn4(self.conv4(output)))
    output = F.relu(self.bn5(self.conv5(output)))
    output = self.adaptive_pool(output)  # 新增自适应池化
    output = output.view(-1, 24 * 10 * 10)
    output = self.fc1(output)
    return output

4. 验证模型参数设备一致性

确保model.to(device)执行后,所有模型参数(包括BatchNorm的running_mean、running_var)都已移至目标设备。可通过打印next(model.parameters()).device确认。


内容的提问来源于stack exchange,提问作者dbel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 02:25:41