Skip to content

分布式数据加载与预处理

高效的分布式数据加载是训练吞吐的关键瓶颈。本文介绍分布式场景下的数据加载、预处理和 I/O 优化实践。

分布式数据加载与预处理

分布式采样

确保各 GPU 获取不重叠的数据分片:

python
from torch.utils.data import DataLoader, DistributedSampler

dataset = MyDataset()
sampler = DistributedSampler(
    dataset,
    num_replicas=dist.get_world_size(),
    rank=dist.get_rank(),
    shuffle=True,
    seed=42,
)

dataloader = DataLoader(
    dataset,
    batch_size=per_gpu_batch_size,
    sampler=sampler,
    num_workers=4,
    pin_memory=True,
    persistent_workers=True,
)

# 每 epoch 设置 sampler
for epoch in range(num_epochs):
    sampler.set_epoch(epoch)
    for batch in dataloader:
        train_step(batch)

预处理与 Tokenization

预计算 Tokenization

python
# 预先将文本数据 Tokenize 并保存
from transformers import AutoTokenizer
import numpy as np

tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-7b-hf")

def preprocess_and_save(text_files, output_dir):
    for i, file in enumerate(text_files):
        with open(file, "r") as f:
            text = f.read()

        tokens = tokenizer.encode(text)
        np.save(f"{output_dir}/tokens_{i}.npy", np.array(tokens, dtype=np.int32))

内存映射加载

python
class MMapDataset(Dataset):
    def __init__(self, data_path, seq_length=2048):
        self.data = np.load(data_path, mmap_mode="r")
        self.seq_length = seq_length

    def __len__(self):
        return (len(self.data) - 1) // self.seq_length

    def __getitem__(self, idx):
        start = idx * self.seq_length
        end = start + self.seq_length
        return torch.tensor(self.data[start:end], dtype=torch.long)

数据格式对比

格式读取速度压缩率随机访问推荐
CSV/JSON支持不推荐
Tokenized NPY支持推荐
MDS (Mosaic)最快支持推荐
WebDataset有限大规模推荐
Parquet有限存储/分析

高性能数据管道

python
# 多阶段数据管道
# 1. 预处理:Tokenize + Pack
# 2. 格式转换:NPY / MDS
# 3. 加载:MMap + DistributedSampler
# 4. 预取:DataLoader num_workers + pin_memory

dataloader = DataLoader(
    dataset,
    batch_size=batch_size,
    sampler=sampler,
    num_workers=8,              # 多进程加载
    pin_memory=True,            # 预分配 CUDA 内存
    persistent_workers=True,    # 保持 Worker 进程
    prefetch_factor=4,          # 预取 4 个 batch
)

GPU 利用率检查

如果 GPU 利用率低于 80%,大概率是数据加载瓶颈。检查 nvidia-smi 的 GPU 利用率和 iostat 的磁盘 I/O。增加 num_workers 和使用 MMap 通常能解决问题。

数据一致性

分布式训练中,必须确保各 Rank 使用相同的数据 Shuffle 顺序(通过 set_epoch 和固定 seed),否则训练结果不可复现。

相关资源

最近更新