自定义数据集训练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
相关产品推荐
相关产品推荐

