分布式训练场景下ShardedByS3Key的正确用法及代码示例咨询
分布式训练中使用ShardedByS3Key的实现方案
核心逻辑概述
ShardedByS3Key是AWS SageMaker用于优化分布式训练数据摄入的策略,它会基于S3对象的键值对数据做分片划分,让每个训练节点仅处理分配给自己的专属数据分片,避免节点间重复读取数据,大幅提升数据加载效率。
PyTorch 实现示例及调整要点
1. 完整代码示例
import os import torch import torch.distributed as dist from torch.utils.data import Dataset, DataLoader class CustomDataset(Dataset): def __init__(self, data_files): self.data_files = data_files self.data = self._load_local_data() def _load_local_data(self): # 替换为你的实际数据加载逻辑(如读取CSV、图片、二进制文件等) data = [] for file_path in self.data_files: with open(file_path, 'r', encoding='utf-8') as f: data.extend(f.readlines()) return data def __len__(self): return len(self.data) def __getitem__(self, idx): # 替换为你的数据预处理逻辑 return self.data[idx].strip() def main(): # SageMaker会自动初始化分布式环境,无需手动设置地址端口 dist.init_process_group(backend='nccl') rank = dist.get_rank() world_size = dist.get_world_size() # 获取当前节点分配到的本地数据目录(SageMaker自动下载S3分片到此处) train_data_dir = os.environ['SM_CHANNEL_TRAIN'] # 遍历目录下所有专属数据文件 local_data_files = [os.path.join(train_data_dir, f) for f in os.listdir(train_data_dir)] dataset = CustomDataset(local_data_files) # 关键:无需使用DistributedSampler # ShardedByS3Key已经在数据层面完成节点级分片,每个节点的数据集是独立完整的分片 dataloader = DataLoader(dataset, batch_size=32, shuffle=True, num_workers=4) # 训练逻辑示例 for epoch in range(10): for batch_idx, batch_data in enumerate(dataloader): # 替换为你的模型训练步骤 print(f"Rank {rank}, Epoch {epoch}, Batch {batch_idx}, Data count: {len(batch_data)}") dist.destroy_process_group() if __name__ == "__main__": main()
2. 关键调整说明
- 彻底移除DistributedSampler:因为ShardedByS3Key已经在S3存储层完成了数据的节点级分片,每个节点只会拿到属于自己的那部分数据文件,再用DistributedSampler会导致节点仅处理自身分片的子集,造成数据浪费或训练不充分。
- 直接读取本地目录:SageMaker会自动将当前节点对应的S3分片下载到
SM_CHANNEL_TRAIN环境变量指定的本地目录,直接读取该目录下的文件即可。
TensorFlow 实现示例及调整要点
1. 完整代码示例
import os import tensorflow as tf def main(): # 获取当前节点分配到的本地数据目录 train_data_dir = os.environ['SM_CHANNEL_TRAIN'] # 匹配目录下所有数据文件(根据数据格式调整通配符) local_data_files = tf.io.gfile.glob(os.path.join(train_data_dir, "*.txt")) # 构建tf.data流水线 dataset = tf.data.TextLineDataset(local_data_files) # 常规数据处理步骤:打乱、分批、预取 dataset = dataset.shuffle(buffer_size=10000) dataset = dataset.batch(32) dataset = dataset.prefetch(tf.data.AUTOTUNE) # 模型定义示例 model = tf.keras.Sequential([ tf.keras.layers.TextVectorization(max_tokens=1000), tf.keras.layers.Dense(64, activation='relu'), tf.keras.layers.Dense(1, activation='sigmoid') ]) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy']) # 启动训练 model.fit(dataset, epochs=10) if __name__ == "__main__": main()
2. 关键调整说明
- 无需额外分布式数据划分:不用调用
tf.distribute.Strategy的experimental_distribute_dataset等方法,ShardedByS3Key已经完成了节点级的数据分片,每个节点的tf.data只需处理本地专属数据即可。 - 直接基于本地文件构建Dataset:依赖SageMaker自动完成的S3分片下载,直接读取
SM_CHANNEL_TRAIN目录下的文件构建数据集即可。
SageMaker 训练任务配置示例
提交训练任务时,需在数据源配置中指定ShardedByS3Key作为分片策略:
from sagemaker.pytorch import PyTorch from sagemaker.inputs import TrainingInput # 定义训练数据输入 train_input = TrainingInput( s3_data="s3://your-bucket/path-to-training-data/", distribution="ShardedByS3Key" ) # 初始化PyTorch训练器(TensorFlow训练器用法类似) estimator = PyTorch( entry_point="train.py", role="your-sagemaker-iam-role", instance_count=2, # 分布式训练节点数量 instance_type="ml.p3.2xlarge", framework_version="2.1.0", py_version="py39" ) # 启动训练任务 estimator.fit({"train": train_input})
内容的提问来源于stack exchange,提问作者Philipp Schmid
相关产品推荐
相关产品推荐

