如何在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
相关产品推荐
相关产品推荐

