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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 16:36:04