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

分布式训练场景下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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 02:25:20