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

PyTorch Dataset返回结构异常:自定义数据集批量输出不符预期

自定义图文Dataset的DataLoader批次结构问题排查与解决

问题根源

你遇到的是PyTorch DataLoader默认collate_fn的行为导致的结构不符:
默认的collate逻辑会将batch中每个样本对应位置的元素拼接——比如每个样本返回的image是长度为3的列表,默认逻辑会把所有样本的第1张图拼成一组,第2张图拼成一组,以此类推,最终得到3个包含4张图的元组;标签label同理,会被拼成3个包含4个标签的元组。而你需要的是保留每个样本的image列表和label列表作为独立单元,组成对应batch_size长度的集合。

同时注意原代码存在一个逻辑错误:循环中你用处理后的张量img和文件名self.gold_images[idx]比较,这永远不会相等,需要修改为比较文件名。

解决方案

1. 修正Dataset中的逻辑错误

class ImageTextDataset(Dataset):
    def __init__(self, data_dir, train_df, tokenizer, feature_extractor, data_type, device, text_augmentation=False):
        self.data_dir = data_dir
        
        if data_type == "train":
            self.tokenizer = tokenizer
            self.feature_extractor=feature_extractor
            self.transforms = transforms.Compose([
                transforms.Resize([512,512]),
                transforms.ToTensor(), 
                transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
            ])
            self.all_image_names = list(train_df['images'])
            self.keywords = list(train_df['word'])
            self.context = list(train_df['description'])
            self.gold_images = list(train_df['gold_image'])

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

    def __getitem__(self, idx):
        context = self.context[idx]
        keyword = self.keywords[idx]
        
        label = []
        image_paths = self.all_image_names[idx]
        images = []
        # 修改变量名避免覆盖,用文件名和gold_image比较
        for img_name in image_paths:
            path = os.path.join(self.data_dir, "trial_images_v1", img_name)
            img = Image.open(path)
            
            if img.mode != "RGB":
                img = img.convert('RGB')
            img_tensor = self.transforms(img)
            images.append(img_tensor)
            label.append(1.0) if img_name == self.gold_images[idx] else label.append(0.0)

        return (context, images, label)

2. 自定义collate_fn实现期望的批次结构

def custom_collate_fn(batch):
    # batch是列表,每个元素是Dataset返回的(context, images, label)元组
    contexts = [item[0] for item in batch]
    image_lists = [item[1] for item in batch]
    label_lists = [item[2] for item in batch]
    # 按需求返回元组或列表形式
    return tuple(contexts), image_lists, label_lists

3. 创建DataLoader时指定自定义collate_fn

# 假设已初始化dataset实例
dataloader = DataLoader(
    dataset,
    batch_size=4,
    shuffle=True,
    collate_fn=custom_collate_fn
)

效果验证

此时每个batch的结构会完全符合你的期望:

  • 第一个元素:('word_1', 'word_2', 'word_3', 'word_4')
  • 第二个元素:[[image_1,image_2,image_3],[image_4,image_5,image_6],[image_7,image_8,image_9], [image_10,image11,image_12]]
  • 第三个元素:[[0,0,1],[1,0,0],[0,1,0],[0,1,0]]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 23:45:29