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

SageMaker Pipeline处理步骤能否使用ImageGenerator实现图像分类

结论

完全可以在SageMaker Pipeline的Processing Step里用Keras的ImageDataGenerator,没有任何功能层面的限制,核心只要对齐Processing运行环境、数据挂载路径和你原有Notebook的代码逻辑即可。

实现思路
  • 选择匹配的Processing运行镜像
    直接用AWS预置的SageMaker TensorFlow Processing镜像即可,镜像已经预装了对应版本TensorFlow、Pillow等ImageDataGenerator的运行依赖,不需要自己从头构建镜像。注意镜像带的TensorFlow版本要和后续训练步骤用的版本保持一致,避免版本兼容问题。

  • 配置Processing Step的输入输出路径映射
    Processing Step运行在独立的计算实例上,不会直接读取Notebook本地文件,你需要把存原始图像的S3路径挂载到容器内的本地目录,比如将训练集S3路径挂载到/opt/ml/processing/input/train、验证集挂载到/opt/ml/processing/input/val,这个路径就是后续传给flow_from_directory的data_dir参数,和你原有代码的路径逻辑完全兼容。

    注意:路径映射不要写错,你代码里自带的os.path.exists(data_dir)校验就是卡这一步的,路径不对会直接抛出找不到图像资源的报错。

  • 封装Processing入口脚本
    你现有的load_data、get_flow_from_directory函数可以直接复用,不需要修改核心逻辑,只需要补几部分内容:

    1. 用argparse接收Processing Step传入的输入路径、批次大小、图像宽高等参数,不要把HEIGHT、WIDTH、batch_size这类参数硬编码在脚本里
    2. 如果要在Processing阶段做离线数据增强、把增强后的图像存下来给后续训练步骤用,拿到generator之后,迭代生成批次数据,把增强后的图像按类别目录结构存到Processing的输出路径(比如/opt/ml/processing/output/train_aug),SageMaker会自动把这个路径下的所有文件同步回你指定的S3位置
    3. 如果只是要在Processing阶段做数据集校验、生成类别映射文件,直接调用get_flow_from_directory拿到class_indices,把映射存成json放到输出路径即可,后续训练步骤可以直接读取这个映射,避免训练、推理阶段类别顺序不一致

    适配后的入口脚本参考:

    import os
    import argparse
    import json
    from tensorflow.keras.preprocessing.image import ImageDataGenerator
    
    def load_data(mode):
        if mode == 'TRAIN':
            datagen = ImageDataGenerator(
                rescale=1. / 255,
                rotation_range = 0.5,
                shear_range=0.2,
                zoom_range=0.2,
                width_shift_range = 0.2,
                height_shift_range = 0.2,
                fill_mode = 'nearest',
                horizontal_flip=True)
        else:
            datagen = ImageDataGenerator(rescale=1. / 255)
        return datagen
    
    
    def get_flow_from_directory(datagen,
                                data_dir,
                                batch_size,
                                img_height,
                                img_width,
                                shuffle=True):
        assert os.path.exists(data_dir), ("Unable to find images resources for input")
        generator = datagen.flow_from_directory(data_dir,
                                                class_mode = "categorical",
                                                target_size=(img_height, img_width),
                                                batch_size=batch_size,
                                                shuffle=shuffle
                                                )
        print('Labels are: ', generator.class_indices)
        return generator
    
    if __name__ == "__main__":
        parser = argparse.ArgumentParser()
        parser.add_argument("--train-input-path", type=str, default="/opt/ml/processing/input/train")
        parser.add_argument("--val-input-path", type=str, default="/opt/ml/processing/input/val")
        parser.add_argument("--output-path", type=str, default="/opt/ml/processing/output")
        parser.add_argument("--batch-size", type=int, default=32)
        parser.add_argument("--img-height", type=int, default=224)
        parser.add_argument("--img-width", type=int, default=224)
        args = parser.parse_args()
    
        # 加载训练集generator,校验数据集
        train_datagen = load_data("TRAIN")
        train_gen = get_flow_from_directory(
            train_datagen,
            args.train_input_path,
            args.batch_size,
            args.img_height,
            args.img_width
        )
        # 保存类别映射
        os.makedirs(args.output_path, exist_ok=True)
        with open(os.path.join(args.output_path, "class_indices.json"), "w") as f:
            json.dump(train_gen.class_indices, f, indent=2)
        
        # 如果需要做离线增强存图,在这里迭代train_gen保存增强后的图像到输出路径即可
    
  • 组装Pipeline步骤
    用FrameworkProcessor初始化TensorFlow处理类,传入对应版本号、执行角色、实例规格,配置好输入S3路径对应的ProcessingInput、输出S3路径对应的ProcessingOutput,指定写好的入口脚本,就可以把这个步骤插到Pipeline的数据准备节点位置。

常见踩坑
  • 没必要把训练时的在线增强逻辑放到Processing Step:ImageDataGenerator的实时增强本来就是训练过程中在线执行的,如果不需要提前生成固定的增强数据集,这部分逻辑直接留在Training Step的训练脚本里就行,放到Processing Step只会增加Pipeline运行时长和存储成本。
  • 不要随意修改输出文件的权限:Processing容器默认用root用户运行,要是手动改了输出文件权限,后续训练步骤读取时可能报权限错误,保持默认权限即可。
  • 大数据集不要全量加载到内存:flow_from_directory本身是按批次读取数据,不要把整个数据集一次性迭代完存在内存里,选的Processing实例内存只要能放下单批次数据即可,否则容易触发OOM。

内容的提问来源于stack exchange,提问作者Cindy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 19:39:38