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
相关产品推荐
相关产品推荐

