SageMaker Pipe模式对接S3目录TFRecord时fit调用无报错挂起问题
问题根因与修复方案
你遇到的无限挂起问题是PipeModeDataset的特性和训练逻辑不匹配导致的,按以下步骤修复即可:
1. 给PipeModeDataset显式添加重复操作
PipeModeDataset默认仅会遍历对应通道下的所有文件一次,当你训练需要跑多个epoch时,第一轮数据读取完成后管道就没有新数据输出,训练进程会无限等待输入导致挂起。你需要给PipeModeDataset显式添加repeat()操作,无参数传入时会无限重复数据,配合训练步数控制停止即可。
2. 明确指定训练/验证步数
当你使用无限重复的数据集时,训练框架无法自动判断每轮epoch的结束节点,你需要在调用model.fit()时手动传入steps_per_epoch和validation_steps参数,取值为对应数据集的总样本数除以batch size即可。
3. 确认依赖版本匹配
你使用的TensorFlow 2.3需要对应版本的sagemaker-tensorflow包支持,在你的source目录下新增requirements.txt文件,添加以下内容避免版本不兼容:
sagemaker-tensorflow==2.3.0
修正后的数据集配置代码示例
import tensorflow as tf if __name__ == "__main__": arg_parser = argparse.ArgumentParser() arg_parser.add_argument("--batch-size", type=int, default=1) arg_parser.add_argument("--pipe_mode", type=int, default=0) # 新增参数用于指定步数,也可以直接在代码里计算总样本数得到 arg_parser.add_argument("--steps_per_epoch", type=int, default=1000) arg_parser.add_argument("--validation_steps", type=int, default=100) arg_parser.add_argument("--train_dir", type=str, default=os.environ.get("SM_CHANNEL_TRAINING")) arg_parser.add_argument( "--validation_dir", type=str, default=os.environ.get("SM_CHANNEL_VALIDATION") ) arg_parser.add_argument("--model_dir", type=str) args, _ = arg_parser.parse_known_args() AUTOTUNE = tf.data.experimental.AUTOTUNE if args.pipe_mode == 1: from sagemaker_tensorflow import PipeModeDataset train_ds = PipeModeDataset(channel="training", record_format='TFRecord') val_ds = PipeModeDataset(channel="validation", record_format='TFRecord') # 新增repeat操作 train_ds = train_ds.repeat() val_ds = val_ds.repeat() else: train_files = tf.data.Dataset.list_files(args.train_dir + '/*tfrecord') val_files = tf.data.Dataset.list_files(args.validation_dir + '/*tfrecord') train_ds = tf.data.TFRecordDataset(filenames=train_files, num_parallel_reads=AUTOTUNE) val_ds = tf.data.TFRecordDataset(filenames=val_files, num_parallel_reads=AUTOTUNE) train_ds = ( train_ds.map(tfrecord_parser, num_parallel_calls=AUTOTUNE) .batch(args.batch_size) .prefetch(AUTOTUNE) ) val_ds = ( val_ds.map(tfrecord_parser, num_parallel_calls=AUTOTUNE) .batch(args.batch_size) .prefetch(AUTOTUNE) ) # 训练时传入步数参数 model.fit( train_ds, steps_per_epoch=args.steps_per_epoch, validation_data=val_ds, validation_steps=args.validation_steps, epochs=10 )
补充说明
你当前的channel参数配置是正确的,PipeModeDataset的channel参数值和你调用fit()时传入的s3_data_channels字典的键一一对应即可,不需要额外修改这部分配置。
内容的提问来源于stack exchange,提问作者FFT
相关产品推荐
相关产品推荐

