如何让DataLoader返回BoundingBoxes类型而非普通张量?
让DataLoader返回BoundingBoxes类型的实现方法
问题原因
DataLoader默认的collate_fn会将BoundingBoxes这类TVTensor对象转换为普通张量,丢失其格式、画布尺寸等关键元数据;直接使用tv_tensors.set_return_type("TVTensor")会触发错误,因为样本中的category是int类型,无法被转成TVTensor子类。
标准解决方案:自定义collate_fn
通过实现自定义的collate函数,单独处理BoundingBoxes类型的数据,保留其元数据,同时正常处理其他类型字段:
1. 实现自定义collate函数
import torch from torchvision.tv_tensors import BoundingBoxes def custom_collate(batch): collated = {} keys = batch[0].keys() for key in keys: samples = [item[key] for item in batch] if isinstance(samples[0], BoundingBoxes): # 堆叠BoundingBoxes的张量部分 stacked_tensor = torch.stack([bbox.as_subclass(torch.Tensor) for bbox in samples]) # 重新构建BoundingBoxes对象,保留元数据 collated[key] = BoundingBoxes( stacked_tensor, format=samples[0].format, canvas_size=samples[0].canvas_size ) else: # 其他类型用默认堆叠逻辑 collated[key] = torch.utils.data.default_collate(samples) return collated
2. 在DataLoader中使用该函数
dataloader = torch.utils.data.DataLoader( dataset, batch_size=2, collate_fn=custom_collate )
3. 验证效果
遍历DataLoader时,box_coord和box_tight会保持BoundingBoxes类型,同时保留元数据:
for batch in dataloader: print(type(batch['box_coord'])) # 输出 <class 'torchvision.tv_tensors.bounding_boxes.BoundingBoxes'> print(batch['box_coord'].format) # 输出 BoundingBoxFormat.XYXY print(batch['box_coord'].canvas_size) # 输出 (416, 416)
补充说明
- 自定义collate_fn的核心是手动保留BoundingBoxes的元数据,默认collate仅处理张量本身,会丢弃TVTensor的额外属性
tv_tensors.set_return_type更适合在transforms流程中使用,确保transform输出TVTensor,但无法处理非张量类型字段(如int类型的category),因此单独用它解决DataLoader的batch问题会触发错误
内容的提问来源于stack exchange,提问作者Wang
相关产品推荐
相关产品推荐

