如何从PyTorch DataLoader中获取未被拆分的List类型返回数据?
问题原因
- DataLoader默认使用内置的
default_collate函数处理单样本拼接为batch的逻辑,该函数会递归遍历所有序列类型(列表、元组等),将不同样本同位置的元素进行拼接。 - 你的Dataset
__getitem__返回的第二个元素是长度为3的列表,default_collate识别到该结构后,会将所有样本的列表第0位、第1位、第2位分别拼接,最终得到你看到的按位置拆分重组的结果。
解决方案
你只需要自定义collate_fn参数,覆盖默认的拼接逻辑即可,示例修改代码如下:
from torch.utils.data import DataLoader,Dataset tests = [('test resume1',[1,2,3]), ('test resume2',['a','b','c']), ('test resume3',['Q',"W",'E']), ('test resume4',[',','.','/']), ('test resume5',['!','@','#'])] # 自定义collate函数,按需要的格式拼接batch def custom_collate(batch): # batch是每个样本__getitem__返回值组成的列表 batch_x = tuple(sample[0] for sample in batch) # 如果需要y的每个元素为元组,改成[tuple(sample[1]) for sample in batch]即可 batch_y = [sample[1] for sample in batch] return [batch_x, batch_y] class testdataset(Dataset): def __init__(self,data): self.x = [item[0] for item in data] self.y = [item[1] for item in data] def __getitem__(self,index): return self.x[index],self.y[index] def __len__(self): return len(self.x) temp = testdataset(tests) print(temp[0]) # DataLoader传入自定义的collate_fn pack = DataLoader(temp,batch_size=2,shuffle=True, collate_fn=custom_collate) for i,unit in enumerate(pack): print(i,type(unit),len(unit)) print(unit)
运行后即可得到你期望的输出格式:
('test resume1', [1, 2, 3]) 0 <class 'list'> 2 [('test resume2', 'test resume4'), [['a', 'b', 'c'], [',', '.', '/']]] 1 <class 'list'> 2 [('test resume5', 'test resume1'), [['!', '@', '#'], [1, 2, 3]]] 2 <class 'list'> 2 [('test resume3',), [['Q', 'W', 'E']]]
内容的提问来源于stack exchange,提问作者liuzw
相关产品推荐
相关产品推荐

