数据加载管线与性能优化
高效的数据加载是训练速度的关键瓶颈。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 会导致内存压力和进程切换开销。