Skip to content

数据加载管线与性能优化

高效的数据加载是训练速度的关键瓶颈。PyTorch 的 DataLoader 提供了多进程加载、预取和批量采样等优化机制。

数据加载管线

核心 API

python
from torch.utils.data import Dataset, DataLoader

class CustomDataset(Dataset):
    def __init__(self, data_path):
        self.data = load_data(data_path)
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        return self.data[idx]

loader = DataLoader(
    dataset,
    batch_size=32,
    shuffle=True,
    num_workers=4,      # 多进程加载
    pin_memory=True,    # CUDA 内存预分配
    prefetch_factor=2,  # 预取批次
    persistent_workers=True  # 保持 Worker 进程
)

num_workers 调优

num_workers 不是越多越好。一般设为 CPU 核心数的一半。过多的 Worker 会导致内存压力和进程切换开销。

相关资源

最近更新