使用Trainer与ConstantLengthDataset的分布式训练问题修复
修复分布式训练下的数据集重复与序列数异常问题
问题根源
- IterableDataset分布式默认行为:每个GPU对应的进程会独立遍历完整数据集,导致不同GPU拿到完全重复的缓冲区数据。
- 无进程级数据分片:自定义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,提问作者имя
相关产品推荐
相关产品推荐

