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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 01:28:09