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

如何在AWS S3上使用Torchvision ImageFolder加载数据集?

解决方案:无需AWS官方插件从S3加载ImageFolder格式数据集

1. 用fsspec+s3fs将S3挂载为本地虚拟文件系统

借助fsspec和s3fs可以把S3存储桶映射成本地可访问的路径,让Torchvision的ImageFolder直接读取,完全复用原有本地加载逻辑。

  • 先安装依赖:
pip install fsspec s3fs
  • 代码示例:
import fsspec
from torchvision.datasets import ImageFolder
from torch.utils.data import DataLoader

# 建立S3连接,挂载指定桶到虚拟路径
fs = fsspec.filesystem('s3')
fs.mount('s3://your-target-bucket', '/mnt/s3-virtual')

# 像读取本地文件夹一样使用ImageFolder
dataset = ImageFolder(root='/mnt/s3-virtual/your-dataset-path')
dataloader = DataLoader(dataset, batch_size=32)

前提:已配置AWS凭证(环境变量、~/.aws/credentials文件等)

2. 自定义Dataset类实现S3图像加载

继承PyTorch的Dataset基类,自己解析S3上的目录结构(模拟ImageFolder的标签-目录映射逻辑),并实现图像字节流的读取与转换。

代码示例:

import boto3
from torch.utils.data import Dataset
from torchvision import transforms
from PIL import Image
import io

class S3ImageFolder(Dataset):
    def __init__(self, s3_bucket, s3_prefix, transform=None):
        self.s3_client = boto3.client('s3')
        self.bucket = s3_bucket
        self.prefix = s3_prefix
        self.transform = transform
        self.image_keys = []
        self.labels = []
        self.label_map = {}

        # 遍历S3目录,构建标签与图像路径的映射
        paginator = self.s3_client.get_paginator('list_objects_v2')
        for page in paginator.paginate(Bucket=s3_bucket, Prefix=s3_prefix):
            for obj in page.get('Contents', []):
                key = obj['Key']
                if key.lower().endswith(('.jpg', '.png', '.jpeg')):
                    # 用图像所在目录名作为标签
                    label_name = key.split('/')[-2]
                    if label_name not in self.label_map:
                        self.label_map[label_name] = len(self.label_map)
                    self.image_keys.append(key)
                    self.labels.append(self.label_map[label_name])

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

    def __getitem__(self, idx):
        key = self.image_keys[idx]
        label = self.labels[idx]
        # 从S3读取图像字节流并转换为PIL格式
        response = self.s3_client.get_object(Bucket=self.bucket, Key=key)
        img_data = response['Body'].read()
        image = Image.open(io.BytesIO(img_data)).convert('RGB')

        if self.transform:
            image = self.transform(image)
        return image, label

# 实例化自定义数据集
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor()
])
dataset = S3ImageFolder(
    s3_bucket='your-bucket-name',
    s3_prefix='dataset-root-path',
    transform=transform
)
dataloader = DataLoader(dataset, batch_size=32, num_workers=4)

3. 基于TorchData的S3数据管道

利用TorchData的S3迭代器组件,先列出S3上的图像文件,再通过映射操作解析图像和标签,适配DataLoader使用。

  • 安装依赖:
pip install torchdata
  • 代码示例:
from torchdata.datapipes.iter import S3FileLister, S3FileLoader
from torchvision import transforms
from PIL import Image
import io
from torch.utils.data import DataLoader

# 列出S3指定路径下的所有图像文件
file_list_dp = S3FileLister('s3://your-bucket/dataset-path', masks=['*.jpg', '*.png'])
# 加载文件字节流
file_load_dp = file_list_dp.read_from_s3()

# 定义解析函数:从文件名提取标签,转换图像格式
def parse_s3_item(file_key, file_data):
    label = file_key.split('/')[-2]
    image = Image.open(io.BytesIO(file_data)).convert('RGB')
    transform = transforms.Compose([
        transforms.Resize((224, 224)),
        transforms.ToTensor()
    ])
    return transform(image), label

# 应用解析逻辑
dataset_dp = file_load_dp.map(parse_s3_item)
# 转换为DataLoader
dataloader = DataLoader(dataset_dp, batch_size=32)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 23:38:17