分布式数据加载与预处理
高效的分布式数据加载是训练吞吐的关键瓶颈。本文介绍分布式场景下的数据加载、预处理和 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),否则训练结果不可复现。