PyTorch DataLoader处理超内存数据集:memmap仍存内存问题求方案
解决超大Numpy数据集的PyTorch DataLoader内存溢出问题
现有代码的核心问题
- 初始化阶段一次性创建所有memmap对象并存储,大量文件的memmap元数据会累积占用内存;遍历memmap获取shape时,系统可能预加载部分数据到内存。
__getitem__中用np.array()强制将memmap切片转为内存数组,加上DataLoader多进程的资源复制,极易触发内存峰值。- 多进程模式下,每个worker会独立初始化Dataset,重复创建memmap进一步加剧内存占用。
解决方案
方案1:延迟创建Memmap对象
不在初始化时预加载所有memmap,仅存储文件路径,在__getitem__中按需打开文件读取数据,同时通过读取npy文件头获取样本数,避免加载整个文件。
修改后的Dataset代码:
import bisect import glob from pathlib import Path import numpy as np import torch class LargeNumpyDataset(torch.utils.data.Dataset): def __init__(self, root: str, dataset: str): # 仅存储文件路径,不提前创建memmap self.list_y_paths = glob.glob(str(Path(root) / "Ground_truth" / dataset / f"Ground_truth_{dataset}_*.npy")) self.list_s2_paths = [p.replace("Ground_truth", "Sentinel-2") for p in self.list_y_paths] # 读取文件头获取样本数,无需加载完整数据 self.start_indices = [0] self.data_count = 0 for y_path in self.list_y_paths: with open(y_path, 'rb') as f: np.lib.format.read_magic(f) header = np.lib.format.read_array_header_1_0(f) self.data_count += header[0] self.start_indices.append(self.data_count) def __getitem__(self, index): # 定位目标文件与内部索引 memmap_idx = bisect.bisect_right(self.start_indices, index) - 1 idx_in_file = index - self.start_indices[memmap_idx] # 按需打开memmap并读取数据,直接转tensor避免拷贝 with np.load(self.list_s2_paths[memmap_idx], mmap_mode='r') as s2_mmap: s2_tensor = torch.as_tensor(s2_mmap[idx_in_file, :, :, :].astype('float32')) with np.load(self.list_y_paths[memmap_idx], mmap_mode='r') as target_mmap: target_tensor = torch.as_tensor(target_mmap[idx_in_file][1]) return {"S2": s2_tensor, "Target": target_tensor} def __len__(self): return self.data_count
方案2:优化DataLoader多进程配置
多进程模式下,每个worker会复制Dataset实例,通过以下设置减少内存浪费:
persistent_workers=True:worker进程在epoch结束后保留,避免重复初始化- 控制
num_workers数量:建议不超过CPU核心数的一半 - 非GPU场景关闭
pin_memory,减少内存拷贝开销
DataLoader初始化示例:
dataset = LargeNumpyDataset(root="/your/root/path", dataset="your_dataset_name") dataloader = torch.utils.data.DataLoader( dataset, batch_size=32, num_workers=4, persistent_workers=True, shuffle=True )
方案3:消除不必要的数据拷贝
- 直接用
torch.as_tensor()从memmap切片创建tensor,跳过np.array()的内存拷贝步骤 - 读取时直接通过
astype()转换数据类型,避免后续额外操作
方案4:大文件拆分(可选)
如果单个npy文件体积远超RAM的1/4,建议将其拆分为更小的文件,降低单个memmap的内存映射开销,同时提升数据读取的并行效率。
内容的提问来源于stack exchange,提问作者Simon Madec
相关产品推荐
相关产品推荐

