如何在AWS SageMaker训练自定义TensorFlow模型时读取外部文件?
这个问题我之前在SageMaker上训练自定义TensorFlow模型时也碰到过!你遇到的IOError本质原因是:Python内置的open()和os.path工具只能处理本地文件系统的路径,没法直接识别S3的s3://格式路径——SageMaker训练容器默认不会自动把你的S3路径挂载到本地,所以直接访问肯定会报错。
下面给你三种可行的解决方案,按推荐程度排序:
方案一:将S3文件作为训练数据挂载到本地(最推荐)
如果这个词汇表文件是训练依赖的核心资源,最好在创建SageMaker训练作业时,通过InputDataConfig把S3路径挂载到容器本地。这样不仅读取效率高,还能避免手动处理S3下载的逻辑。
步骤1:配置训练作业的输入数据
在定义TensorFlow Estimator时,添加输入数据配置,把你的S3词汇表路径挂载到容器本地的指定目录:
from sagemaker.tensorflow import TensorFlow estimator = TensorFlow( entry_point='your_training_script.py', role='your_sagemaker_execution_role', instance_count=1, instance_type='ml.m5.xlarge', framework_version='2.12', py_version='py310', # 配置S3数据挂载 input_data_config=[ { 'DataSource': { 'S3DataSource': { 'S3Uri': 's3://<bucket_name>/data/<prefix>/', 'S3DataType': 'S3Prefix', 'S3DataDistributionType': 'FullyReplicated' } }, 'TargetAttributeName': 'vocab', # 自定义挂载目录的名称 'ContentType': 'application/octet-stream', 'InputMode': 'File' } ] ) # 启动训练 estimator.fit()
步骤2:在model_fn中读取本地挂载的文件
SageMaker会把你配置的S3数据自动放到容器的/opt/ml/input/data/<TargetAttributeName>路径下,直接用open()读取即可:
import os import pickle def model_fn(features, labels, mode, params): # 拼接本地挂载的词汇表路径 local_vocab_path = os.path.join('/opt/ml/input/data/vocab', 'vocab.pkl') # 直接读取本地文件 with open(local_vocab_path, 'rb') as f: vocab = pickle.load(f) n_vocab = len(vocab) # 后续的模型构建逻辑...
方案二:使用s3fs库直接读取S3文件
如果你不想修改训练作业的配置,可以用s3fs库——它能让你像操作本地文件一样读写S3资源,非常方便。
步骤1:添加依赖
在你的训练脚本的requirements.txt里加上s3fs(如果没有这个文件,就新建一个放在脚本同目录):
s3fs>=2023.6.0
步骤2:在model_fn中用s3fs读取文件
直接用s3fs.S3FileSystem().open()替代原生的open():
import s3fs import pickle def model_fn(features, labels, mode, params): vocab_s3_path = 's3://<bucket_name>/data/<prefix>/vocab.pkl' # 初始化S3文件系统客户端 fs = s3fs.S3FileSystem() # 像读本地文件一样读取S3文件 with fs.open(vocab_s3_path, 'rb') as f: vocab = pickle.load(f) n_vocab = len(vocab) # 后续模型逻辑...
方案三:用boto3手动下载到临时目录
这是最基础的方法,手动用boto3把S3文件下载到容器的临时目录(/tmp是SageMaker容器里的可写目录),再读取本地文件。
import boto3 import os import pickle from botocore.exceptions import ClientError def model_fn(features, labels, mode, params): # 拆分S3路径为桶名和文件键 bucket_name = '<bucket_name>' s3_vocab_key = 'data/<prefix>/vocab.pkl' # 本地临时路径 local_vocab_path = '/tmp/vocab.pkl' # 初始化S3客户端 s3 = boto3.client('s3') try: # 下载文件到本地临时目录 s3.download_file(bucket_name, s3_vocab_key, local_vocab_path) # 读取本地文件 with open(local_vocab_path, 'rb') as f: vocab = pickle.load(f) n_vocab = len(vocab) # 后续模型逻辑... except ClientError as e: # 捕获S3访问错误,方便调试 raise Exception(f"Failed to download vocab file from S3: {str(e)}")
重要提醒:权限检查
不管用哪种方法,都要确保你的SageMaker训练角色(就是创建Estimator时指定的role)有s3:GetObject权限访问目标S3桶和文件,否则会出现权限拒绝的错误。
内容的提问来源于stack exchange,提问作者Dimitris Poulopoulos

