如何配置PyTorch DataLoader处理每行返回多输出的变换函数
PyTorch DataLoader 一行生成多数据点的处理方案
问题背景
现有myfile.csv内容:
imagefile,label train/0/16585.png,0 train/0/56789.png,0
需要构建PyTorch DataLoader,遍历后返回两倍于输入行数的数据点,示例输出:
>>> dp = MyDataPipe(csvfile) >>> for row in dp.train_dataloader: ... print(row) ... (tensor([1.23, 4.56, 7.89]), 0) (tensor([9.87, 6.54, 3.21]), 1) (tensor([9.99, 8.88, 7.77]), 0) (tensor([1.11, 2.22, 9.87]), 1)
遇到两个问题:
- 变换函数
optimus_prime每行用yield返回2组数据时,DataLoader无法正确拆分这些数据点 - 使用两个变换函数分别生成数据时,第二个函数触发
TypeError: tuple indices must be integers or slice not str
解决方案
1. 处理一行生成多数据点的问题
PyTorch DataLoader默认会把变换函数的返回值直接作为一个数据元素,所以如果变换函数用yield返回多个结果,需要用展开操作把生成器的输出拆成单个数据点。
方法1:用yield from展开生成器
直接在迭代逻辑中用yield from遍历变换函数的生成器输出,将每个结果单独作为数据点返回:
from torch.utils.data import DataLoader, IterableDataset import pandas as pd import torch def optimus_prime(row): # 模拟生成两个数据点:原标签和原标签+1 label = row['label'] # 模拟加载图像生成张量 tensor1 = torch.randn(3) yield (tensor1, label) tensor2 = torch.randn(3) yield (tensor2, label + 1) class MyDataPipe(IterableDataset): def __init__(self, csv_path): self.df = pd.read_csv(csv_path) def __iter__(self): for _, row in self.df.iterrows(): # 展开生成器的每个输出 yield from optimus_prime(row) @property def train_dataloader(self): return DataLoader(self, batch_size=None) # 逐个返回数据点,不做batch合并
方法2:返回列表后逐个yield
如果变换函数返回包含多组数据的列表,直接遍历列表逐个yield:
def optimus_prime(row): label = row['label'] tensor1 = torch.randn(3) tensor2 = torch.randn(3) return [(tensor1, label), (tensor2, label + 1)] class MyDataPipe(IterableDataset): def __init__(self, csv_path): self.df = pd.read_csv(csv_path) def __iter__(self): for _, row in self.df.iterrows(): for item in optimus_prime(row): yield item
2. 解决TypeError: tuple indices must be integers or slice not str错误
这个错误的核心原因是:第一个变换函数将原csv行的字典结构转成了tuple,但第二个变换函数仍然尝试用字符串索引(如row['imagefile'])访问数据。
修复方式1:按索引访问tuple数据
如果第一个变换返回tuple,第二个变换要按索引提取值:
def transform1(row): # 将csv行字典转成tuple格式 return (row['imagefile'], row['label']) def transform2(row): # 按索引获取tuple中的值,而非字符串索引 img_path, label = row tensor1 = torch.randn(3) yield (tensor1, label) tensor2 = torch.randn(3) yield (tensor2, label + 1) class MyDataPipe(IterableDataset): def __init__(self, csv_path): self.df = pd.read_csv(csv_path) def __iter__(self): for _, row in self.df.iterrows(): transformed = transform1(row) yield from transform2(transformed)
修复方式2:保持字典结构
如果要避免索引方式的变化,让第一个变换函数返回字典而非tuple:
def transform1(row): # 保持原字典结构,后续变换仍可用字符串索引 return {'imagefile': row['imagefile'], 'label': row['label']}
内容的提问来源于stack exchange,提问作者alvas
相关产品推荐
相关产品推荐

