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

在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里的图片,完全不用手动下载数据,还能按需加载文件(适合大数据量场景)。

步骤:

  1. 确保依赖包已安装(SageMaker的默认conda环境可能已包含,若没有则执行):
!pip install sagemaker s3fs
  1. 用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)
  1. 直接使用挂载后的本地路径调用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的封装功能)。

步骤:

  1. 安装TensorFlow IO:
!pip install tensorflow-io
  1. 配置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
  1. 构建自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 14:02:39