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

如何让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 00:31:03