PyTorch Dataset内存泄漏问题求助(GPU无泄漏)
PyTorch DataLoader内存泄漏问题分析
问题根源
- 多进程内存复制开销
当DataLoader设置num_workers>0时,PyTorch会创建对应数量的子进程加载数据。每个子进程都会完整复制父进程中的Dataset对象——包括你在__init__里创建的、包含1000万条数据的file_paths大列表。你设置了num_workers=10,相当于额外有10个进程各自持有一份1000万条的列表,直接导致RAM占用剧增。
当shuffle=True时,DataLoader会提前生成一份包含所有样本索引的随机排列列表,这个大列表同样会被每个子进程复制,所以内存占用增长速度会更快。
- 子进程内存回收机制缺陷
默认情况下,DataLoader会复用worker进程(persistent_workers=True),这些子进程中复制的大对象不会在epoch结束后被及时释放,持续占用内存,表现为“泄漏”。
解决方案
- 延迟/惰性加载大列表
不要在Dataset的__init__中一次性生成file_paths,改为在__getitem__中动态生成单条路径,避免父进程的大对象被多进程重复复制。示例修改:
from torch.utils.data import Dataset import torch class MemLeakDataset(Dataset): def __init__(self): self.total_samples = 10_000_000 def __len__(self): return self.total_samples def __getitem__(self, idx): label = idx image = [] return image, label
调整worker数量
根据机器内存情况降低num_workers,减少子进程复制的内存开销。使用共享内存存储大对象
通过torch.multiprocessing的共享内存机制让所有worker共享同一份file_paths,避免重复复制。示例:
from torch.utils.data import Dataset import torch from multiprocessing import Manager class MemLeakDataset(Dataset): def __init__(self, file_paths): self.file_paths = file_paths def __len__(self): return len(self.file_paths) def __getitem__(self, idx): label = self.file_paths[idx][1] image = [] return image, label if __name__ == "__main__": with Manager() as manager: file_paths = manager.list([(f"fake_file_path_{i}", i) for i in range(10_000_000)]) train_dataset = MemLeakDataset(file_paths) train_dataloader = torch.utils.data.DataLoader(train_dataset, batch_size=256, shuffle=False, drop_last=False, num_workers=10) count = 0 for _ in train_dataloader: count += 1 print(count)
- 关闭worker复用
设置persistent_workers=False,让每个epoch结束后销毁worker进程释放内存,但会增加每次epoch的启动时间开销。
内容的提问来源于stack exchange,提问作者Anvar Ganiev
相关产品推荐
相关产品推荐

