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

使用Trainer与ConstantLengthDataset的分布式训练问题修复

修复分布式训练下的数据集重复与序列数异常问题

问题根源

  1. IterableDataset分布式默认行为:每个GPU对应的进程会独立遍历完整数据集,导致不同GPU拿到完全重复的缓冲区数据。
  2. 无进程级数据分片:自定义Dataset未针对分布式环境做数据分片处理,所有进程共享同一份数据源迭代逻辑。

具体修复步骤

1. 初始化时注入分布式环境信息

在ConstantLengthDataset的__init__方法中添加分布式相关参数,自动获取或传入当前进程的rank和总进程数(world_size):

import os

class ConstantLengthDataset(IterableDataset):
    def __init__(
        self, 
        tokenizer, 
        dataset, 
        infinite=False,
        seq_length=8192,
        num_of_sequences=1024,
        chars_per_token=3.6,
        rank=None,
        world_size=None,
    ):
        self.tokenizer = tokenizer
        self.concat_token_id = tokenizer.eos_token_id
        self.dataset = dataset
        self.seq_length = seq_length
        self.infinite = infinite
        self.current_size = 0
        self.max_buffer_size = seq_length * chars_per_token * num_of_sequences
        # 从环境变量或传入参数获取分布式信息,兼容单机/分布式场景
        self.rank = rank if rank is not None else int(os.environ.get("RANK", 0))
        self.world_size = world_size if world_size is not None else int(os.environ.get("WORLD_SIZE", 1))

2. 实现进程级数据分片逻辑

修改__iter__方法,让每个进程只遍历属于自己的数据集分片,避免跨进程数据重复:

def __iter__(self):
    iterator = iter(self.dataset)
    more_examples = True
    
    # 初始化时跳过当前rank之前的元素,定位到分片起始位置
    for _ in range(self.rank):
        try:
            next(iterator)
        except StopIteration:
            if not self.infinite:
                more_examples = False
                break

    while more_examples:
        buffer, buffer_len = [], 0
        while True:
            if buffer_len >= self.max_buffer_size:
                break
            try:
                # 按world_size步长获取元素,每个进程只取属于自己分片的数据
                for _ in range(self.world_size - 1):
                    next(iterator)
                item = next(iterator)
                buffer.append(item["content"])
                buffer_len += len(buffer[-1])
            except StopIteration:
                if self.infinite:
                    iterator = iter(self.dataset)
                    # 重置数据集后重新定位到当前分片的起始位置
                    for _ in range(self.rank):
                        next(iterator)
                else:
                    more_examples = False
                    break

        tokenized_inputs = self.tokenizer(buffer, truncation=False)["input_ids"]
        all_token_ids = []
        for tokenized_input in tokenized_inputs:
            all_token_ids.extend(tokenized_input + [self.concat_token_id])
        
        for i in range(0, len(all_token_ids), self.seq_length):
            input_ids = all_token_ids[i : i + self.seq_length]
            if len(input_ids) == self.seq_length:
                self.current_size += 1
                yield {
                    "input_ids": torch.LongTensor(input_ids),
                    "labels": torch.LongTensor(input_ids),
                }

3. 初始化Dataset时传入分布式参数

创建ConstantLengthDataset实例时,传入分布式环境的rank和world_size(也可依赖环境变量自动获取):

from torch.distributed import get_rank, get_world_size

train_dataset = ConstantLengthDataset(
    tokenizer=your_tokenizer,
    dataset=your_raw_dataset,
    infinite=True,
    seq_length=8192,
    rank=get_rank(),
    world_size=get_world_size(),
)

4. 调整Trainer配置确保批次符合预期

  • 保持per_device_train_batch_size=1,Trainer会自动计算总批次(总批次=单GPU批次×GPU数量)。
  • 将dataloader_num_workers设为0或1,避免多worker场景下的额外数据重复(多worker需额外处理worker级分片,单worker更易控制)。
  • 若需减少每个GPU单次处理的序列数,可降低num_of_sequences参数,缩小缓冲区生成的序列总量。

内容的提问来源于stack exchange,提问作者имя

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 08:27:28