在Amazon SageMaker中如何将AWS S3存储桶数据传入tf.keras.utils.image_dataset_from_directory()的img_dir变量
解决方案:在SageMaker中让
tf.keras.utils.image_dataset_from_directory读取S3数据 我在SageMaker环境里处理过类似的需求,分享几个能让你不用手动下载S3数据,直接给tf.keras.utils.image_dataset_from_directory喂数据的可行方案:
方案1:将S3存储桶挂载为本地文件系统(推荐)
SageMaker支持通过FUSE将S3桶挂载成本地目录,这样image_dataset_from_directory就能像访问本地路径一样读取S3里的图片,完全不用手动下载数据,还能按需加载文件(适合大数据量场景)。
步骤:
- 确保依赖包已安装(SageMaker的默认conda环境可能已包含,若没有则执行):
!pip install sagemaker s3fs
- 用SageMaker Session挂载S3路径到本地:
from sagemaker import Session import os # 初始化SageMaker会话 sagemaker_session = Session() # 配置你的S3路径和本地挂载目录 s3_bucket = "data-ma5852" s3_image_prefix = "image_data" local_mount_path = "/home/ec2-user/SageMaker/mounted_s3_images" # 创建本地挂载目录(不存在则自动创建) os.makedirs(local_mount_path, exist_ok=True) # 挂载S3路径到本地目录 sagemaker_session.fs.mount(f"s3://{s3_bucket}/{s3_image_prefix}", local_mount_path)
- 直接使用挂载后的本地路径调用
image_dataset_from_directory:
import tensorflow as tf train_ds = tf.keras.utils.image_dataset_from_directory( local_mount_path, validation_split=0.2, subset="training", seed=123, image_size=(224, 224) # 替换为你的目标图片尺寸 ) val_ds = tf.keras.utils.image_dataset_from_directory( local_mount_path, validation_split=0.2, subset="validation", seed=123, image_size=(224, 224) )
方案2:用TensorFlow IO直接读取S3数据(自定义数据集流程)
如果不想挂载文件系统,也可以通过TensorFlow IO扩展直接读取S3路径的图片,不过需要自己实现数据解析和划分逻辑(替代image_dataset_from_directory的封装功能)。
步骤:
- 安装TensorFlow IO:
!pip install tensorflow-io
- 配置AWS凭证(SageMaker执行角色已具备S3访问权限的话,这一步可省略):
from sagemaker import Session import os sagemaker_session = Session() credentials = sagemaker_session.boto_session.get_credentials() os.environ['AWS_ACCESS_KEY_ID'] = credentials.access_key os.environ['AWS_SECRET_ACCESS_KEY'] = credentials.secret_key os.environ['AWS_REGION'] = sagemaker_session.boto_region_name
- 构建自定义TF Dataset:
import tensorflow as tf import tensorflow_io as tfio # 定义S3图片路径的通配符(假设你的图片按分类存于子文件夹:s3://bucket/image_data/class1/xxx.jpg) s3_file_pattern = "s3://data-ma5852/image_data/*/*.jpg" # 获取所有图片文件路径 file_dataset = tf.data.Dataset.list_files(s3_file_pattern) # 定义图片解析和标签提取函数 def parse_image(file_path): # 从文件路径中提取分类标签(取倒数第二个文件夹名称) label = tf.strings.split(file_path, os.sep)[-2] # 读取并预处理图片 img = tf.io.read_file(file_path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (224, 224)) img = tf.keras.applications.resnet50.preprocess_input(img) # 可选,根据你的模型调整预处理 return img, label # 映射解析函数到数据集 raw_dataset = file_dataset.map(parse_image, num_parallel_calls=tf.data.AUTOTUNE) # 划分训练集和验证集(按8:2比例) dataset_size = tf.data.experimental.cardinality(raw_dataset).numpy() train_size = int(0.8 * dataset_size) train_ds = raw_dataset.take(train_size).shuffle(1000).batch(32).prefetch(tf.data.AUTOTUNE) val_ds = raw_dataset.skip(train_size).batch(32).prefetch(tf.data.AUTOTUNE)
方案3:手动下载S3数据到本地(你已实现的备选)
这个是你当前在使用的方法,适合数据量较小的场景,用SageMaker的S3Downloader工具一键同步数据:
from sagemaker.s3 import S3Downloader local_img_dir = "/home/ec2-user/SageMaker/image_data" S3Downloader.download("s3://data-ma5852/image_data/", local_img_dir)
内容的提问来源于stack exchange,提问作者riot12
相关产品推荐
相关产品推荐

