如何用boto3将S3目录转为os.Path格式适配Torchvision ImageFolder?
问题描述
我想在Torchvision的ImageFolder中使用存储在AWS S3桶里的数据集,但ImageFolder要求传入类os.Path的对象,而用boto3只能拿到目录或文件列表,满足不了要求。
当前用boto3加载数据集的代码:
# Initialize S3 client s3 = boto3.client('s3') response_test = s3.list_objects_v2(Bucket=BUCKET_NAME_test, Prefix=PREFIX_test) jpg_files_test = [item['Key'] for item in response_test.get('Contents', []) if item['Key'].lower().endswith('.jpeg')]
用Torchvision加载数据的代码:
transform_train = transforms.Compose([ transforms.RandomResizedCrop(args.input_size, scale=(0.2, 1.0), interpolation=3), # 3 is bicubic transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean=[0.10231429, 0.18161748, 0.26542443], std=[0.06512465, 0.10374635, 0.16232868])]) dataset_train = datasets.ImageFolder(os.path.join(args.data_path, 'train'), transform=transform_train) # 'train' print(dataset_train)
运行后报错:
TypeError: expected str, bytes or os.PathLike object, not list
请问怎么通过boto3把S3目录转为os.Path格式,或者用类似策略解决这个问题?
解决方案
方法一:用s3fs将S3桶映射为本地可访问路径
s3fs可以让S3路径被Python文件系统API识别,让ImageFolder直接兼容S3存储。
- 安装依赖:
pip install s3fs
- 修改代码直接使用S3路径:
import s3fs # 自动读取AWS凭证(环境变量、~/.aws/credentials等) fs = s3fs.S3FileSystem() # 直接传入S3路径给ImageFolder dataset_train = datasets.ImageFolder('s3://你的桶名/train', transform=transform_train)
确保你的AWS账号拥有该S3桶的读取权限。
方法二:自定义Dataset类替代ImageFolder
如果不想依赖第三方库,可基于boto3实现一个模拟ImageFolder行为的自定义数据集类:
import torch from torch.utils.data import Dataset from PIL import Image import io import boto3 class S3ImageFolder(Dataset): def __init__(self, bucket_name, prefix, transform=None): self.s3 = boto3.client('s3') self.bucket = bucket_name self.prefix = prefix self.transform = transform self.file_list = [] self.class_to_idx = {} # 识别所有类别目录(S3结构需为:prefix/类别名/图片文件) dir_response = self.s3.list_objects_v2(Bucket=bucket_name, Prefix=prefix, Delimiter='/') classes = [cp['Prefix'].split('/')[-2] for cp in dir_response.get('CommonPrefixes', [])] self.class_to_idx = {cls: idx for idx, cls in enumerate(classes)} # 遍历每个类别,收集所有图片文件 for cls in classes: cls_prefix = f"{prefix}{cls}/" file_response = self.s3.list_objects_v2(Bucket=bucket_name, Prefix=cls_prefix) for item in file_response.get('Contents', []): if item['Key'].lower().endswith(('.jpeg', '.jpg', '.png')): self.file_list.append((item['Key'], self.class_to_idx[cls])) def __len__(self): return len(self.file_list) def __getitem__(self, idx): key, label = self.file_list[idx] # 从S3读取图片到内存 obj = self.s3.get_object(Bucket=self.bucket, Key=key) image = Image.open(io.BytesIO(obj['Body'].read())).convert('RGB') if self.transform: image = self.transform(image) return image, label # 使用自定义数据集 dataset_train = S3ImageFolder(BUCKET_NAME_test, PREFIX_test, transform=transform_train)
方法三:直接使用Torchvision原生S3支持(Torchvision ≥0.13)
从Torchvision 0.13版本开始,部分数据集类已原生支持S3路径,无需额外配置(前提是AWS凭证已正确设置):
dataset_train = datasets.ImageFolder('s3://你的桶名/train', transform=transform_train)
底层仍依赖s3fs,但无需显式初始化。
内容的提问来源于stack exchange,提问作者AlternativeWaltz
相关产品推荐
相关产品推荐

