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

如何从含图像路径与标签的Pandas DataFrame加载数据到PyTorch DataLoader?

问题:PyTorch是否有类似Keras flow_from_dataframe()的函数处理DataFrame图像数据集?

我计划使用PyTorch完成图像二分类任务,训练数据以CSV文件存储,包含img_path、printer_id、print_id、has_under_extrusion(标签值为0或1)列,CSV格式示例如下:

img_path,printer_id,print_id,has_under_extrusion
101/1678589738/1678589914.060332.jpg,101,1678589738,1
101/1678589738/1678589914.462857.jpg,101,1678589738,1
101/1678589738/1678589914.875075.jpg,101,1678589738,1
...

我已经用pd.read_csv()将数据导入Pandas DataFrame,希望基于这个DataFrame构建PyTorch的torch.utils.data.Dataset并封装为DataLoader,同时对每张图像应用变换。想知道PyTorch是否提供类似Keras中flow_from_dataframe()的函数,还是需要自行实现?


回答:

PyTorch核心库没有直接对应flow_from_dataframe()的开箱即用函数,但有两种成熟方案可以实现需求:

方案1:自行实现自定义Dataset(推荐,灵活可控)

这是PyTorch处理此类场景的标准实践,步骤清晰且完全可控。只需继承torch.utils.data.Dataset,重写核心方法即可完成数据加载、变换和标签匹配。

示例代码:

import torch
import os
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
import pandas as pd
from PIL import Image

class ImageBinaryDataset(Dataset):
    def __init__(self, dataframe, root_dir, transform=None):
        self.df = dataframe
        self.root_dir = root_dir  # 图像文件的根路径,与df中的img_path拼接得到完整路径
        self.transform = transform
        self.labels = dataframe['has_under_extrusion'].values

    def __len__(self):
        return len(self.df)

    def __getitem__(self, idx):
        # 拼接完整图像路径
        img_full_path = os.path.join(self.root_dir, self.df.iloc[idx]['img_path'])
        # 加载图像并转为RGB格式
        image = Image.open(img_full_path).convert('RGB')
        # 获取对应标签
        label = self.labels[idx]
        # 应用图像变换
        if self.transform:
            image = self.transform(image)
        return image, label

# 定义图像预处理变换
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])

# 初始化数据集与DataLoader
# 假设你的DataFrame名为df,图像根目录为'./dataset_images'
dataset = ImageBinaryDataset(df, root_dir='./dataset_images', transform=transform)
dataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4)

方案2:使用第三方库的现成工具

如果不想手动编写Dataset,部分第三方库提供了类似封装:

  • torchvision:新版本的torchvision.datasets模块有一些可扩展的类,但灵活性不如自定义Dataset;
  • pytorch-lightning:其LightningDataModule可以简化数据流程,但底层仍依赖自定义Dataset的逻辑;
  • albumentations:专注于图像增强,可配合自定义Dataset使用,但并非直接替代flow_from_dataframe()的工具。

总体来说,自定义Dataset是最直接、灵活的选择,也是PyTorch生态的标准做法,建议优先采用。

内容的提问来源于stack exchange,提问作者Abe Mankavil

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 22:02:26