自定义PyTorch Dataset提取标签后训练遇batch_size不匹配错误求助
数据集背景
拥有35类图像数据集,所有图像存于同一文件夹,图像名称包含标签信息。示例图像名:D34_Samsung_GalaxyS3Mini-images-flat-D01_I_flat_0001.jpg,对应标签为D01,索引应为34。
最初的Dataset实现及问题
最初实现的Dataset类代码如下:
class MyDataset(Dataset): def __init__(self, imgs , transform = None): self.imgs = imgs self.transform = transform or transforms.ToTensor() self.class_to_idx = {} def __getitem__(self, index): image_path = self.imgs[index] target = image_path.split('_')[0] target = re.findall(r'D\d+.+' , target) image = Image.open(image_path) if self.transform is not None: image = self.transform(image) if target[0] in self.class_to_idx : target = [self.class_to_idx[target[0]]] else : self.class_to_idx[target[0]] = len(self.class_to_idx) target = [self.class_to_idx[target[0]]] return image , target def __len__(self): return len(self.imgs)
测试发现问题:标签始终在0-15区间(与batch_size=16对应),且每次运行图像标签可能不同。原因是class_to_idx在__getitem__中动态生成,标签分配依赖图像访问顺序,且无法覆盖全部35类。
修改后的Dataset实现及训练错误
修改后的代码尝试直接提取标签,但训练时触发错误:
class MyDataset(Dataset): def __init__(self, imgs , transform = None): self.imgs = imgs self.transform = transform or transforms.ToTensor() self.class_to_idx = {} def __getitem__(self, index): image_path = self.imgs[index] target = image_path.split('_')[0] target = target.split('D')[1] target = int(target) image = Image.open(image_path) if self.transform is not None: image = self.transform(image) return image , target def __len__(self): return len(self.imgs)
训练时错误信息:
ValueError: Expected input batch_size (16) to match target batch_size (0).
原模型使用第一个代码可正常训练,排除模型本身问题。
解决建议
1. 修正标签提取逻辑
当前修改后的代码提取的是文件名开头的D34,但真实标签是文件名中的D01,属于标签提取错误。需用正则表达式提取正确的标签字段:
# 在__getitem__中替换标签提取代码 import re # 提取所有D+数字的字段,取第二个即为真实标签D01 target_str = re.findall(r'D\d+', image_path)[1] # 提取标签中的数字 target_num = int(target_str.split('D')[1])
2. 预先构建固定的类-索引映射
第一个代码的动态标签分配导致标签不稳定,需在__init__中预先构建所有类的固定映射(根据你的数据集规则调整,示例为D01到D35对应索引0到34,若D01需对应索引34,自行修改映射逻辑):
def __init__(self, imgs, transform=None): self.imgs = imgs self.transform = transform or transforms.ToTensor() # 预先构建固定的类到索引的映射 self.class_to_idx = {f"D{i:02d}": idx for idx, i in enumerate(range(1, 36))} # 若D01对应索引34,可改为: # self.class_to_idx = {f"D{i:02d}": 35 - i for i in range(1, 36)}
3. 统一标签返回格式并修复维度匹配
第一个代码返回的标签是列表格式(如[34]),修改后的代码返回单个整数,导致DataLoader打包后的张量维度与模型期望不匹配。同时,原模型可能依赖class_to_idx的长度确定输出类别数,空字典会导致模型输出维度为0,触发批量不匹配错误。
修正后的完整__getitem__示例:
def __getitem__(self, index): image_path = self.imgs[index] # 提取正确的标签字符串 target_str = re.findall(r'D\d+', image_path)[1] # 获取对应的索引 target_idx = self.class_to_idx[target_str] image = Image.open(image_path) if self.transform is not None: image = self.transform(image) # 保持与原代码一致的列表格式,或根据模型要求返回单个整数 return image, [target_idx]
4. 验证数据集与模型的维度匹配
训练前打印数据集的样本标签和模型输出形状,确认:
- 标签的维度与模型损失函数要求一致(如CrossEntropyLoss接受一维标签,而MSELoss可能接受二维标签)
- 模型最后一层的输出类别数等于
len(self.class_to_idx)(即35)
内容的提问来源于stack exchange,提问作者Nima

