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

如何配置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)

遇到两个问题:

  1. 变换函数optimus_prime每行用yield返回2组数据时,DataLoader无法正确拆分这些数据点
  2. 使用两个变换函数分别生成数据时,第二个函数触发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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 18:20:54