tf.keras.Model.save_weights在SageMaker TensorFlow Estimator中存S3失败
SageMaker中TensorFlow Estimator直接保存模型权重到S3失败的原因及解决办法
问题背景
在SageMaker Studio中通过TensorFlow Estimator运行训练脚本时,两种保存模型权重的方式里,直接写入S3的方式(model.save_weights('s3://<some-bucket>/...'))在更换新AWS角色后突然失效,抛出错误:
File system scheme 's3' not implemented (file: 's3://<some-bucket>/<some-folder>/directly-saved/model_%Y%m%d%H%M%S')
补充信息:
- 在Notebook中直接执行相同的保存命令无报错
- 在Estimator脚本中使用boto3可以正常完成S3写入
原因分析
- 新角色的权限与信任策略缺失
- 新角色未被授权访问目标S3桶的读写权限,导致训练实例无法向S3写入数据
- 角色的信任策略未允许SageMaker服务扮演该角色,使得训练集群无法获取合法身份凭证访问AWS服务
- 训练容器与Notebook环境的依赖差异
SageMaker Studio Notebook环境默认预装了s3fs、tensorflow-io等支持S3文件系统的依赖,而TensorFlow Estimator的训练容器可能未默认配置这些依赖,导致TensorFlow无法识别s3://协议。之前能正常运行是因为旧角色的权限允许容器间接获取相关配置,或旧角色关联的训练环境已预装依赖。
解决方案
方案1:修复新角色的IAM配置
1.1 添加S3桶访问权限
给新角色添加以下IAM策略,授予目标S3桶的读写权限:
{ "Version": "2012-10-17", "Statement": [ { "Effect": "Allow", "Action": [ "s3:PutObject", "s3:GetObject", "s3:ListBucket" ], "Resource": [ "arn:aws:s3:::<your-bucket-name>", "arn:aws:s3:::<your-bucket-name>/*" ] } ] }
1.2 配置信任策略
确保角色的信任策略允许SageMaker服务扮演该角色:
{ "Version": "2012-10-17", "Statement": [ { "Effect": "Allow", "Principal": { "Service": "sagemaker.amazonaws.com" }, "Action": "sts:AssumeRole" } ] }
方案2:在训练环境中添加S3依赖支持
给TensorFlow Estimator的训练环境安装支持S3的依赖:
- 创建
requirements.txt文件,添加以下内容:
s3fs>=2023.6.0 tensorflow-io>=0.32.0
- 定义TensorFlow Estimator时指定依赖文件:
from sagemaker.tensorflow import TensorFlow estimator = TensorFlow( entry_point='your-training-script.py', role='your-new-role-arn', instance_count=1, instance_type='ml.m5.xlarge', framework_version='2.13', py_version='py310', requirements_file='requirements.txt' )
方案3:沿用boto3替代直接写入
既然boto3方式可正常工作,可继续使用该方案:先将模型权重保存到训练实例临时目录,再通过boto3上传到S3:
import boto3 import os import time # 保存到本地临时目录 local_weights_path = f'/tmp/model_weights_{time.strftime("%Y%m%d%H%M%S", time.gmtime())}' model.save_weights(local_weights_path, save_format='tf') # boto3上传到S3 s3 = boto3.client('s3') s3.upload_file( local_weights_path, '<some-bucket>', f'<some-folder>/directly-saved/model_weights_{time.strftime("%Y%m%d%H%M%S", time.gmtime())}' ) # 清理临时文件 os.remove(local_weights_path)
内容的提问来源于stack exchange,提问作者Andrey Shumilov
相关产品推荐
相关产品推荐

