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

自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 16:55:22