文章

训练数据流水线设计

训练数据流水线设计

设计端到端的高吞吐训练数据流水线,确保 GPU 永远不因数据加载而空转。覆盖流水线各环节的设计决策、参数调优和故障诊断。

1. 数据流水线阶段

S3 / Lustre (原始数据)


预处理节点 (CPU/GPU 集群)
    │  解压 / 转码 / 增强 / 打包

NVMe 本地缓存 (热数据)
    │  rsync 预热后常驻

PyTorch DataLoader (多进程)
    │  prefetch + pin_memory

GPU HBM (训练)

各阶段吞吐量基准

阶段典型吞吐瓶颈类型
S3 读取5-20 GB/s网络带宽
Lustre 读取50-200 GB/sOST 数量/网络
NVMe 本地读7+ GB/s/盘PCIe 带宽
DataLoader 消费1-5 GB/s (CPU 解码)CPU 核数/解码速度
GPU 消费由模型决定训练带宽需求

核心公式:

DataLoader 吞吐 ≥ GPU 消费速度 × 1.2
GPU 消费速度 ≈ (batch_size × sample_size × 2) / step_time

📖 已掌握


2. PyTorch DataLoader 调优

核心参数详解

from torch.utils.data import DataLoader

loader = DataLoader(
    dataset,
    batch_size=256,
    num_workers=8,              # 并行进程数,建议 per-GPU
    prefetch_factor=4,          # 每个 worker 预取 batch 数
    pin_memory=True,            # 锁页内存 → GPU 传输更快
    pin_memory_device="cuda",   # PyTorch 2.0+,指定目标设备
    persistent_workers=True,    # 复用 worker 进程,避免 fork 开销
    drop_last=True,             # 丢弃不完整 batch,利于多卡对齐
)

参数调优指南

# num_workers 调优:从 CPU 核数 / GPU 数开始,逐步增加直到 CPU 利用率饱和
# 经验值:4-16 workers per GPU,取决于数据预处理复杂度

# 快速诊断:观察 DataLoader 迭代时间
python -c "
import time
for i, batch in enumerate(loader):
    t = time.time()
    # ...训练 step...
    print(f'Data wait: {t - last:.3f}s')  # < 0.1s 为正常
    last = time.time()
"
参数推荐值说明
num_workersmin(16, n_cpu // n_gpu)太大导致 CPU 争抢,太小 GPU 等待
prefetch_factor2-4增加 GPU 端缓冲深度
pin_memoryTrue约 2x 加速 Host→Device 传输
persistent_workersTrue避免每次 epoch fork,节省 1-3s/epoch
multiprocessing_context"forkserver"在某些环境下比 fork 更稳定

常见瓶颈诊断

# 1. 查看 DataLoader 进程状态
htop -p $(pgrep -f "torch" | head -20)

# 2. 使用 PyTorch Profiler 定位预处理瓶颈
# 在 training loop 中:
with torch.profiler.profile(
    activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
    schedule=torch.profiler.schedule(wait=1, warmup=1, active=3),
) as prof:
    for batch in loader:
        # training step
        prof.step()

📖 已掌握


3. 数据格式优化

格式对比

格式读取方式随机访问压缩适用场景
WebDatasettar 分片 + URL索引✅ 分片级✅ gzip/zstd大规模图文/视频
Mosaic StreamingDataset自定义流式✅ 样本级可控流式训练
TFRecordprotobuf 序列化✅ 样本级✅ gzipTF 生态
原始文件直接读取小数据集
HDF5分层归档科学计算/NLP

WebDataset 实战

# 制作 WebDataset 分片
tar -cf dataset-000.tar --sort=name /data/samples/

# PyTorch 加载
import webdataset as wds

dataset = (
    wds.WebDataset("shards/dataset-{000000..000999}.tar")
    .shuffle(1000)
    .decode("pil")
    .to_tuple("jpg", "cls")
)
loader = wds.WebLoader(dataset, batch_size=256, num_workers=8)

Mosaic StreamingDataset

from streaming import StreamingDataset

class CustomDataset(StreamingDataset):
    def __init__(self, local, remote, **kwargs):
        super().__init__(local=local, remote=remote, **kwargs)

    def __getitem__(self, idx):
        obj = super().__getitem__(idx)
        return transform(obj['image']), obj['label']

# 自动管理本地缓存和远端拉取
dataset = CustomDataset(local="/mnt/nvme/cache", remote="s3://bucket/dataset")

📖 已掌握


4. 本地 NVMe 缓存策略

缓存决策矩阵

场景是否缓存到 NVMe理由
数据集 < 本地 NVMe 容量✅ 全量缓存消除网络延迟
数据集 >> NVMe 容量⚠️ 按需缓存热点使用 streaming 格式
多 epoch 训练✅ 缓存减少重复网络读取
单遍训练❌ 直接流式读取缓存无收益
多作业共享数据集✅ 缓存(只读)避免 Lustre 热点

rsync 预热脚本

#!/bin/bash
# warmup-dataset.sh — 训练前将数据集从 Lustre 同步到本地 NVMe

LUSTRE_PATH="/mnt/lustre/datasets/${DATASET_NAME}"
NVME_PATH="/mnt/nvme/datasets/${DATASET_NAME}"

echo "[$(date)] 开始预热 ${DATASET_NAME}..."

# 并行 rsync,按子目录拆分加速
find "${LUSTRE_PATH}" -maxdepth 1 -type d | \
  parallel -j 8 "rsync -avP --progress {} ${NVME_PATH}/"

# 校验完整性
diff <(ls -R "${LUSTRE_PATH}" | md5sum) <(ls -R "${NVME_PATH}" | md5sum)
echo "[$(date)] 预热完成"

缓存生命周期管理

训练前:
  → 检查 NVMe 剩余空间
  → 清理过期缓存 (find -mtime +7 -delete)
  → rsync 预热最新数据集

训练中:
  → 只读挂载 (/mnt/nvme),避免误删
  → 监控 NVMe 温度和 SMART 健康状态

训练后:
  → 保留缓存 N 天(高频数据集保留更久)
  → LRU 清理:自动删除最久未访问的数据集

NVMe 性能监控

# 查看 NVMe 读写吞吐
iostat -xmt 1 nvme0n1

# 检查 NVMe 健康状态
nvme smart-log /dev/nvme0
nvme list

# RAID0 条带化(多盘合并吞吐)
mdadm --create /dev/md0 --level=0 --raid-devices=4 \
  /dev/nvme0n1 /dev/nvme1n1 /dev/nvme2n1 /dev/nvme3n1
mkfs.xfs -f /dev/md0
mount -o noatime,nodiratime /dev/md0 /mnt/nvme

📖 已掌握


5. 数据集预分片

分布式训练分片策略

# 自定义 DistributedSampler 替代默认
from torch.utils.data.distributed import DistributedSampler

sampler = DistributedSampler(
    dataset,
    num_replicas=world_size,
    rank=rank,
    shuffle=True,
    drop_last=True,      # 对齐所有 rank 的 batch 数
    seed=42,             # 固定种子,可复现
)

loader = DataLoader(dataset, batch_size=batch_size, sampler=sampler)

预分片文件布局

datasets/
└── imagenet/
    ├── train/
    │   ├── shard_00.tar   # 分配给 rank 0,8,16...
    │   ├── shard_01.tar   # 分配给 rank 1,9,17...
    │   ├── shard_02.tar   # ...
    │   └── shard_NN.tar
    └── val/
        └── val.tar

WebDataset 按 rank 分片

# 每个 rank 独立消费自己的分片
shard_pattern = f"shards/shard_{rank:02d}-{world_size:02d}-*.tar"
dataset = wds.WebDataset(shard_pattern, nodesplitter=wds.split_by_worker)

预分片脚本

#!/bin/bash
# 将原始数据均匀分片到 N 个 tar 文件
N_SHARDS=${1:-128}  # 分片数(建议 128-512)
INPUT_DIR=${2:-"./data"}
OUTPUT_DIR=${3:-"./shards"}

mkdir -p "${OUTPUT_DIR}"

# 列出文件并均分
find "${INPUT_DIR}" -type f | shuf | split -n l/${N_SHARDS} --numeric-suffixes=1 \
  --additional-suffix=".list" - "${OUTPUT_DIR}/files_"

# 按列表打包
for i in $(seq -w 1 ${N_SHARDS}); do
    tar -cf "${OUTPUT_DIR}/shard_${i}.tar" \
        -T "${OUTPUT_DIR}/files_${i}.list" &
done
wait
echo "分片完成: ${N_SHARDS} shards in ${OUTPUT_DIR}"

📖 已掌握


6. 检查点策略

两种策略对比

策略写入路径训练阻塞数据安全实施复杂度
同步写共享存储GPU → Lustre✅ 阻塞 step✅ 高
异步分层写GPU → NVMe → Lustre❌ 不阻塞⚠️ 仅 NVMe 副本

推荐:异步分层检查点

import torch
import threading
import subprocess
from pathlib import Path

class AsyncCheckpointer:
    """先写本地 NVMe,后台线程异步同步到 Lustre/S3"""

    def __init__(self, local_dir, remote_dir, keep_latest=3):
        self.local = Path(local_dir)
        self.remote = Path(remote_dir)
        self.keep = keep_latest
        self.local.mkdir(parents=True, exist_ok=True)
        self._pending_syncs = []

    def save(self, model, optimizer, step, metrics=None):
        ckpt_path = self.local / f"ckpt_step{step:08d}.pt"
        torch.save({
            "step": step,
            "model_state_dict": model.state_dict(),
            "optimizer_state_dict": optimizer.state_dict(),
            "metrics": metrics,
        }, ckpt_path)

        # 异步同步,不阻塞训练
        t = threading.Thread(target=self._sync, args=(ckpt_path, step))
        t.start()
        self._pending_syncs.append(t)

    def _sync(self, local_path, step):
        dest = self.remote / local_path.name
        subprocess.run(["rsync", "-aP", str(local_path), str(dest)])
        self._cleanup_old()

    def _cleanup_old(self):
        """保留最近 N 个检查点"""
        ckpts = sorted(self.local.glob("ckpt_step*.pt"))
        for old in ckpts[:-self.keep]:
            old.unlink()

# 使用
ckpt = AsyncCheckpointer("/mnt/nvme/checkpoints", "/mnt/lustre/checkpoints")
# 每 1000 step 保存一次
if step % 1000 == 0:
    ckpt.save(model, optimizer, step)

多机检查点协调

import torch.distributed as dist

def save_distributed_checkpoint(model, optimizer, step):
    # 使用 PyTorch 分布式保存(>= 2.0)
    from torch.distributed.checkpoint import save

    state_dict = {
        "model": model.state_dict(),
        "optimizer": optimizer.state_dict(),
    }
    # 先写到本地 NVMe
    save(state_dict, checkpoint_id=f"/mnt/nvme/ckpts/step_{step}")
    dist.barrier()  # 等待所有 rank

    # rank 0 负责同步到 Lustre
    if dist.get_rank() == 0:
        subprocess.run([
            "rsync", "-aP",
            "/mnt/nvme/ckpts/",
            "/mnt/lustre/checkpoints/"
        ])

📖 已掌握


7. 数据 Stall 诊断

诊断信号

信号含义阈值
GPU SM 利用率 < 80%GPU 等待数据< 80% 持续 > 2s
CPU I/O Wait 高存储或网络瓶颈> 10%
DataLoader 迭代时间 > GPU step 时间数据供给不足ratio > 1.0
GPU 功耗偏低GPU 未满载< 80% TDP

诊断命令

# 1. 实时监控 GPU SM 利用率
nvidia-smi dmon -s pucv -d 2

# 2. 监控 DataLoader 端 CPU 负载
mpstat -P ALL 1

# 3. 监控 IO 等待
iostat -xmt 1 nvme0n1

# 4. 使用 torch.profiler 精确测量
# 在训练代码中:
from torch.profiler import profile, ProfilerActivity

with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
             record_shapes=True,
             with_stack=True) as prof:
    for step, batch in enumerate(loader):
        train_step(batch)
        if step > 20:
            break

# Chrome trace 分析
prof.export_chrome_trace("trace.json")
# 在 chrome://tracing 打开,查看 DataLoader 和 GPU 时间线重叠度

常见 Stall 原因与修复

问题: DataLoader CPU 进程占用 100%,GPU 仍在等待
原因: 数据解码/增强是瓶颈
修复:
  → 使用 DALI (NVIDIA Data Loading Library) GPU 解码
  → 离线预处理为预解码格式 (numpy/pt)
  → 减少 per-sample augmentation,改用 batch-level

问题: GPU SM 利用率间歇性掉到 0
原因: num_workers 太小或 persistent_workers=False
修复:
  → num_workers += 4,开启 persistent_workers=True
  → prefetch_factor 调至 4

问题: 多 epoch 后期 GPU 利用率下降
原因: shuffle buffer 耗尽或 epoch 切换时有 stall
修复:
  → 使用 IterableDataset + 循环缓冲区
  → 预取下一个 epoch 数据(双缓冲)

问题: Lustre 读取延迟波动大
原因: 多作业同时读取产生 I/O 争抢
修复:
  → 训练前预热数据集到 NVMe
  → 使用 WebDataset 减少小文件元数据压力

📖 已掌握


实用命令速查

# 数据预热
rsync -avP --info=progress2 /mnt/lustre/dataset/ /mnt/nvme/dataset/

# NVMe 性能测试
fio --name=randread --ioengine=libaio --direct=1 --bs=1M \
    --numjobs=4 --iodepth=64 --rw=randread --runtime=30 \
    --filename=/dev/nvme0n1

# 模拟 GPU 端数据消费速度
python -c "
data = torch.randn(256, 3, 224, 224)
t = time.time()
for _ in range(100):
    data = data.cuda()
    data = data * 2
print(f'{100 / (time.time()-t):.0f} samples/s')
"

# 数据集大小统计
du -sh /mnt/lustre/datasets/*
find /mnt/lustre/datasets -type f | wc -l

关联知识

学习时间

阶段时间备注
骨架创建2026-06-30框架搭建
深度填充2026-06-30七节核心内容 + 实战命令

状态标记

📖 已掌握 — 数据流水线架构、DataLoader 参数调优、数据格式对比、NVMe 缓存生命周期、预分片、异步检查点、Stall 诊断

📝 待补充 — DALI (NVIDIA Data Loading Library) GPU 解码实战、大规模数据集(PB 级)缓存淘汰算法、对象存储 S3 Select 下推优化、跨数据中心数据同步方案