DataLoader分组caption异常:预期与实际输出不符该如何解决?
问题原因
PyTorch DataLoader默认的collate_fn会按返回值的元素位置批量拼接数据。你的__getitem__返回的caps是长度为20的列表,当batch_size=32时,默认逻辑会把32个样本中每个caps的第1个元素、第2个元素……第20个元素分别打包,最终得到20个长度为32的元组,和你预期的32个长度为20的caption列表结构完全相反。
解决方案
自定义collate_fn函数,手动按照样本维度整理batch数据,覆盖默认的拼接逻辑:
1. 编写自定义collate函数
import torch def custom_collate_fn(batch): # batch是一个列表,每个元素对应__getitem__返回的(imgs, caps, cap_len, cls_id, key) # 批量拼接图像张量 imgs = torch.stack([item[0] for item in batch], dim=0) # 按样本收集caption,得到32个长度为20的列表 caps = [item[1] for item in batch] # 整理caption长度、类别ID为张量 cap_len = torch.tensor([item[2] for item in batch], dtype=torch.long) cls_id = torch.tensor([item[3] for item in batch], dtype=torch.long) # 收集图片key keys = [item[4] for item in batch] return imgs, caps, cap_len, cls_id, keys
2. 创建DataLoader时指定collate_fn
from torch.utils.data import DataLoader # 假设你的数据集实例是dataset dataloader = DataLoader( dataset, batch_size=32, shuffle=True, collate_fn=custom_collate_fn # 传入自定义的拼接函数 )
补充说明
如果后续需要将caption字符串转换为可输入模型的数值张量(比如使用词嵌入层),可以在custom_collate_fn中进一步处理:提前准备好词表,把每个字符串替换成对应的索引,再转成形状为(32, 20)的LongTensor,这样更方便模型处理。
内容的提问来源于stack exchange,提问作者hanugm
相关产品推荐
相关产品推荐

