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

如何在AWS SageMaker训练自定义TensorFlow模型时读取外部文件?

解决SageMaker model_fn中读取S3文件的IOError问题

这个问题我之前在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 02:24:24