默认collate_fn报错“expected sequence”是因为其内部对非张量序列(如原始list)调用torch.stack失败,而变长数据无法堆叠;根本原因是default_collate仅支持同形tensor或numpy数组的自动拼接,对嵌套list、混合类型或长度不一的序列束手无策。

为什么默认collate_fn会报错“expected sequence”?
PyTorch的DataLoader默认用default_collate,它假设每个样本的张量形状一致。遇到文本、点云或变长序列(如不同长度的list、torch.Tensor)时,会直接抛出TypeError: expected sequence或RuntimeError: stack expects each tensor to be equal size。
根本原因是default_collate对list尝试调用torch.stack,而变长数据无法堆叠。
- 常见触发场景:NLP中句子长度不一、CV中目标检测的bbox数量动态变化、语音中音频帧数不同
- 关键判断点:如果你的
__getitem__返回类似{"text": [12, 45, 67], "label": 1}这种含原始list的字典,就一定需要自定义 - 注意:
default_collate能处理numpy.ndarray或torch.Tensor的同形拼接,但对嵌套list或混合类型束手无策
如何写一个安全的padding collate_fn?
核心思路是:先提取同字段数据 → 手动pad → 转tensor → 组装batch dict。不要依赖default_collate自动推导。
以NLP为例,假设__getitem__返回{"input_ids": [101, 234, 567], "label": 0}:
def collate_fn(batch):
# 提取所有input_ids,找最大长度
max_len = max(len(item["input_ids"]) for item in batch)
# 手动pad并转tensor
input_ids = [item["input_ids"] + [0] * (max_len - len(item["input_ids"])) for item in batch]
input_ids = torch.tensor(input_ids, dtype=torch.long)
# label直接stack(标量)
labels = torch.tensor([item["label"] for item in batch])
return {"input_ids": input_ids, "label": labels}
- padding值选0还是
tokenizer.pad_token_id取决于任务,别硬写死0 - 别在collate里做tokenize——耗时且破坏DataLoader多进程加速,应在
__getitem__里完成 - 如果字段含None(如部分样本无bbox),需先过滤或统一占位,否则
len(None)报错
处理嵌套结构(如目标检测中的boxes)怎么办?
当一个样本含多个可变长度子项(例如{"image": tensor, "boxes": [[x1,y1,x2,y2], ...], "labels": [1,2]}),不能简单pad到同一长度——box数量差异太大时,会浪费显存且影响loss计算。
- 推荐方案:保持
boxes和labels为list of tensor,batch内各元素独立,只对image做常规stack - 示例返回:
{"image": torch.stack(imgs), "boxes": [torch.tensor(b) for b in boxes_batch], "labels": [torch.tensor(l) for l in labels_batch]} - 后续模型需适配这种输入(如
FasterRCNN的forward就接受list of tensors) - 若必须pad成统一shape(如某些YOLO变种),用
torch.nn.utils.rnn.pad_sequence配合batch_first=True,但注意它要求输入是tensor list,不是普通list
性能陷阱:collate_fn里哪些操作会拖慢DataLoader?
collate_fn运行在主进程(即使num_workers>0),任何重计算或IO都会成为瓶颈。
- 禁止在collate里调用
torchvision.transforms——这些应放在__getitem__里,由worker并行执行 - 避免深拷贝大对象:
copy.deepcopy(batch)比原地构造慢数倍 - 慎用
np.array()中间转换:torch tensor创建比numpy array快,优先用torch.tensor(..., dtype=...) - 调试时加
print?会导致多进程输出混乱且阻塞,改用logging或只在num_workers=0时临时加
最易被忽略的是:把本该在__getitem__里做的预处理(比如读图、resize)挪到collate_fn里,等于放弃所有worker并发优势。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











