SageMaker Studio中TensorFlow estimator训练输入下载路径错误问题
问题根因
- SageMaker Python SDK提交训练任务时,默认会将当前工作目录下的所有文件打包为源码包上传至指定S3桶。你将笔记本导出为HTML后,该HTML文件存放在代码运行的同目录下,被自动纳入了上传的源码文件列表。
- 路径拼接逻辑冲突:你同时给
TensorFlowEstimator的bucket参数赋值为my-bucket,又给fit()方法直接传入带S3协议头的bucket_dir = "s3://my-bucket"作为输入路径,SDK内部拼接路径时会自动追加桶前缀,最终重复拼接出s3://my-bucket/s3://my-bucket/这种包含//的非法路径。 - 你看到
SM_CHANNEL_TRAINING环境变量为None,是因为任务在数据下载阶段就因路径错误提前失败,还没有执行到初始化训练环境、注入channel环境变量的步骤。
修复方案
- 移除Estimator构造参数中冗余的
bucket = bucket配置,避免路径重复拼接。 - 在代码运行的工作目录下创建
.sagemakerignore文件,写入*.html规则,让SDK打包源码时自动排除所有HTML文件,避免无关文件被上传。 - 修正
fit()方法的输入传参逻辑,如果要指定S3输入路径,优先使用官方的TrainingInput类明确标识输入通道,示例代码如下:
from sagemaker.inputs import TrainingInput train_input = TrainingInput(s3_data=bucket_dir, content_type="application/npy") history = aws_estimator.fit({"train": train_input})
- 清理Estimator构造参数里的非官方字段:
my_name、log_name、train_data、train_labels这些参数不属于TensorFlow Estimator的官方构造参数,不会被SDK识别,建议全部移到shared_hyperparameters字典中传递,避免参数解析时出现未定义行为。
内容的提问来源于stack exchange,提问作者kwscott
相关产品推荐
相关产品推荐

