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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 23:54:04