如何修改Hugging Face SageMaker训练脚本读取S3上的CSV数据
修改AWS train.py脚本读取S3中的CSV数据以微调DistilBERT
要替换原脚本中load_from_disk的数据集加载逻辑,改用Hugging Face datasets库直接读取S3上的CSV文件,步骤如下:
1. 确保依赖环境
首先确认训练环境已安装必要依赖:
datasets库(Hugging Face官方数据集工具)s3fs库(用于S3文件系统访问,datasets依赖它处理S3路径)
如果是在AWS SageMaker等托管环境,可通过requirements.txt添加依赖:
datasets>=2.0.0 s3fs>=2023.0.0
2. 替换数据集加载代码
原脚本中使用load_from_disk的部分:
from datasets import load_from_disk # 原加载逻辑 dataset = load_from_disk("./local_dataset_path") train_dataset = dataset["train"] test_dataset = dataset["test"]
替换为以下代码,直接读取S3上的CSV文件:
from datasets import load_dataset # 方式1:分别加载单个文件 train_dataset = load_dataset("csv", data_files="s3://your-bucket-name/path/to/train.csv", split="train") test_dataset = load_dataset("csv", data_files="s3://your-bucket-name/path/to/test.csv", split="train") # 方式2:一次性加载多个文件并指定split(更清晰) dataset = load_dataset( "csv", data_files={ "train": "s3://your-bucket-name/path/to/train.csv", "test": "s3://your-bucket-name/path/to/test.csv" } ) train_dataset = dataset["train"] test_dataset = dataset["test"]
3. 权限配置
- 若在AWS托管环境(如SageMaker训练作业)运行:确保训练作业关联的IAM角色拥有S3读权限(可附加
AmazonS3ReadOnlyAccess策略,或更精细的单Bucket访问策略)。 - 若本地运行:提前通过
aws configure配置AWS凭证,确保当前用户有权限访问目标S3 Bucket。
4. 验证与后续处理
加载完成后,可打印数据集样本确认数据正确:
print(train_dataset[0])
后续的文本预处理(如用DistilBERT tokenizer做分词)、模型训练逻辑无需修改——因为load_dataset返回的Dataset对象和load_from_disk返回的结构完全一致,原脚本的map等操作可直接复用。
内容的提问来源于stack exchange,提问作者kyagu
相关产品推荐
相关产品推荐

