如何修改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)
关键说明:
- 递归转换函数:
_convert_to_numpy可以处理不同结构的batch(单个Tensor、元组/列表、字典),确保所有嵌套的Tensor都被转为numpy数组。 - 迭代器重写:__iter__方法不再直接返回原迭代器,而是逐个处理每个batch,转换后通过
yield返回,这样每次迭代得到的都是numpy数组形式的batch。 - 兼容原有参数:__init__方法保留了
**kwargs,可以传入DataLoader的原有参数(如batch_size、shuffle等)。
内容的提问来源于stack exchange,提问作者Mihai.Mehe
相关产品推荐
相关产品推荐

