如何为AWS SageMaker组织分目录存储的图像训练数据集
SageMaker 跨目录图像-标签对数据集传入方案
SageMaker没有强制要求数据集的存储结构,你可以通过以下两种常用方案适配你当前的跨目录存储结构:
方案1:多输入通道配置(无需调整现有存储结构,适配绝大多数场景)
SageMaker Estimator支持同时传入多个独立的S3路径作为不同的训练输入通道,训练实例启动时会自动将每个通道对应的S3文件同步到本地指定目录,你仅需在训练代码中按文件名匹配源图和标签即可。
配置示例
import sagemaker from sagemaker.estimator import Estimator from sagemaker.inputs import TrainingInput # 定义两个独立的输入通道,分别对应源图和标签图的S3路径 training_input_channels = { "imgs": TrainingInput(s3_data="s3://bucket/v1/imgs", s3_data_type="S3Prefix"), "lbls": TrainingInput(s3_data="s3://bucket/v1/lbls", s3_data_type="S3Prefix") } # 初始化你的训练Estimator,替换为你实际使用的框架、实例配置即可 estimator = Estimator( image_uri="<你的训练镜像URI>", # 可直接使用SageMaker官方提供的PyTorch/TensorFlow等框架镜像 role=sagemaker.get_execution_role(), instance_count=1, instance_type="ml.p3.2xlarge", hyperparameters={} # 传入你的自定义超参数 ) # 传入多通道配置启动训练 estimator.fit(training_input_channels)
训练代码读取逻辑
训练启动后,两个通道的文件会分别映射到训练实例的/opt/ml/input/data/imgs和/opt/ml/input/data/lbls路径,也可以通过预置环境变量获取路径,按文件名匹配读取即可:
import os import cv2 # 读取环境变量获取两个通道的本地路径 img_root = os.environ.get("SM_CHANNEL_IMGS") label_root = os.environ.get("SM_CHANNEL_LBLS") for img_filename in os.listdir(img_root): if not img_filename.endswith(".jpg"): continue # 按你的命名规则匹配标签文件名 label_filename = f"label_{img_filename}" img = cv2.imread(os.path.join(img_root, img_filename)) label = cv2.imread(os.path.join(label_root, label_filename), cv2.IMREAD_GRAYSCALE) # 后续正常执行训练逻辑即可
方案2:Augmented Manifest清单配置(适用于样本过滤、自定义匹配规则场景)
如果需要自定义训练样本范围、不需要全量加载两个目录下的所有文件,可以提前生成JSON Lines格式的样本清单,直接指定每一对源图和标签的S3路径,无需调整现有存储结构。
清单文件示例
创建名为train_manifest.json的清单文件,每行对应一个训练样本:
{"source-ref": "s3://bucket/v1/imgs/image1.jpg", "label-ref": "s3://bucket/v1/lbls/label_image1.jpg"} {"source-ref": "s3://bucket/v1/imgs/image2.jpg", "label-ref": "s3://bucket/v1/lbls/label_image2.jpg"}
将清单文件上传到S3路径,例如s3://bucket/v1/train_manifest.json,再配置Estimator输入:
train_input = TrainingInput( s3_data="s3://bucket/v1/train_manifest.json", s3_data_type="AugmentedManifestFile", attribute_names=["source-ref", "label-ref"], input_mode="Pipe" # 可选Pipe模式流式读取,无需全量下载数据到训练实例 ) estimator.fit({"train": train_input})
然后在训练代码中直接读取清单内容即可获取每对样本的S3地址,按需加载即可。
内容的提问来源于stack exchange,提问作者Evan Lalo
相关产品推荐
相关产品推荐

