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

如何修改torch.utils.data.DataLoader使其返回numpy数组而非张量?

让DataLoader返回numpy数组而非Tensor的实现方法

问题:torch.utils.data.DataLoader默认返回torch.Tensor,能否让它返回numpy数组?我尝试子类化DataLoader,但修改__iter__方法后仍返回Tensor,不知道该如何实现转换(比如用my_tensor.numpy())。

我的代码:

class CustomDataLoader(DataLoader):
    def __init__(self, dataset):
        super().__init__(dataset)
    
    def __iter__(self):
        it_ = super().__iter__()
        print( next(it_))
        print(super().__iter__().__dict__)
        return it_ 
        
c = CustomDataLoader(dataset)
next(iter(c))

解决方案:
你之前的代码只是返回了原DataLoader的迭代器,没有对每个batch做转换。要实现目标,需要在__iter__方法中遍历原迭代器的每个batch,将其中的Tensor递归转换为numpy数组后再返回。

修改后的代码:

import torch
from torch.utils.data import DataLoader, Dataset

# 示例测试数据集
class TestDataset(Dataset):
    def __len__(self):
        return 10
    def __getitem__(self, idx):
        # 返回元组形式的样本,模拟常见的输入+标签结构
        return torch.tensor([idx, idx+1]), torch.tensor([idx*2])

class CustomDataLoader(DataLoader):
    def __init__(self, dataset, **kwargs):
        super().__init__(dataset, **kwargs)
    
    def _convert_to_numpy(self, item):
        # 处理单个Tensor
        if isinstance(item, torch.Tensor):
            return item.numpy()
        # 处理列表/元组类型的结构(比如输入+标签的元组)
        elif isinstance(item, (list, tuple)):
            return type(item)(self._convert_to_numpy(i) for i in item)
        # 处理字典类型的结构(比如多输入的字典)
        elif isinstance(item, dict):
            return {k: self._convert_to_numpy(v) for k, v in item.items()}
        # 非Tensor类型直接返回
        else:
            return item
    
    def __iter__(self):
        # 遍历原DataLoader的每个batch
        for batch in super().__iter__():
            # 转换当前batch的所有Tensor为numpy数组
            yield self._convert_to_numpy(batch)

# 测试代码
dataset = TestDataset()
loader = CustomDataLoader(dataset, batch_size=2)

# 查看返回结果
for batch in loader:
    print("Batch内容:", batch)
    print("第一个元素类型:", type(batch[0]), "形状:", batch[0].shape)
    print("第二个元素类型:", type(batch[1]), "形状:", batch[1].shape)

关键说明:

  1. 递归转换函数:_convert_to_numpy可以处理不同结构的batch(单个Tensor、元组/列表、字典),确保所有嵌套的Tensor都被转为numpy数组。
  2. 迭代器重写:__iter__方法不再直接返回原迭代器,而是逐个处理每个batch,转换后通过yield返回,这样每次迭代得到的都是numpy数组形式的batch。
  3. 兼容原有参数:__init__方法保留了**kwargs,可以传入DataLoader的原有参数(如batch_size、shuffle等)。

内容的提问来源于stack exchange,提问作者Mihai.Mehe

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 08:47:18