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函数可以直接复用,不需要修改核心逻辑,只需要补几部分内容:- 用
argparse接收Processing Step传入的输入路径、批次大小、图像宽高等参数,不要把HEIGHT、WIDTH、batch_size这类参数硬编码在脚本里 - 如果要在Processing阶段做离线数据增强、把增强后的图像存下来给后续训练步骤用,拿到generator之后,迭代生成批次数据,把增强后的图像按类别目录结构存到Processing的输出路径(比如
/opt/ml/processing/output/train_aug),SageMaker会自动把这个路径下的所有文件同步回你指定的S3位置 - 如果只是要在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

