You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch Dataset内存泄漏问题求助(GPU无泄漏)

PyTorch DataLoader内存泄漏问题分析

问题根源

  1. 多进程内存复制开销
    当DataLoader设置num_workers>0时,PyTorch会创建对应数量的子进程加载数据。每个子进程都会完整复制父进程中的Dataset对象——包括你在__init__里创建的、包含1000万条数据的file_paths大列表。你设置了num_workers=10,相当于额外有10个进程各自持有一份1000万条的列表,直接导致RAM占用剧增。

当shuffle=True时,DataLoader会提前生成一份包含所有样本索引的随机排列列表,这个大列表同样会被每个子进程复制,所以内存占用增长速度会更快。

  1. 子进程内存回收机制缺陷
    默认情况下,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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.19 02:38:22