PyTorch超大图像数据集训练:洗牌与归一化内存溢出问题求助
大规模PyTorch数据集的内存高效洗牌与归一化方案
一、解决划分前的洗牌问题
核心思路是只对样本索引洗牌,而非直接操作全量数据,原始内存映射数据全程不加载进内存,仅通过索引定位读取样本:
- 生成全局索引数组:创建包含所有样本序号(0到64151)的数组,该数组体积极小(仅约250KB),可直接在内存中完成洗牌。
- 基于洗牌后的索引划分数据集:拆分索引数组得到训练/测试集对应的索引子集,无需触碰原始数据。
- 自定义Dataset类:通过内存映射读取数据,根据索引返回对应样本。
示例代码:
import numpy as np import torch from torch.utils.data import Dataset, Subset, DataLoader class MemMapImageDataset(Dataset): def __init__(self, mmap_file_path, mean=None, std=None): # 内存映射加载数据集,仅建立磁盘映射,不加载全量数据 self.data = np.memmap( mmap_file_path, dtype='float32', mode='r', shape=(64152, 3, 5, 2, 64, 144) ) self.mean = mean self.std = std def __len__(self): return len(self.data) def __getitem__(self, idx): # 仅读取单个样本到内存 sample = self.data[idx].copy() # copy避免内存映射的只读限制 if self.mean is not None and self.std is not None: # 应用归一化 sample = (sample - self.mean.reshape(3, 1, 1, 1, 1)) / self.std.reshape(3, 1, 1, 1, 1) return torch.tensor(sample, dtype=torch.float32) # 生成并洗牌全局索引 total_indices = np.arange(64152) np.random.shuffle(total_indices) # 划分训练/测试集(按8:2比例) train_split = int(0.8 * len(total_indices)) train_indices = total_indices[:train_split] test_indices = total_indices[train_split:] # 初始化数据集 full_dataset = MemMapImageDataset("your_dataset_path.npy") train_dataset = Subset(full_dataset, train_indices) test_dataset = Subset(full_dataset, test_indices) # 构建DataLoader,训练时开启batch级洗牌 train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4) test_loader = DataLoader(test_dataset, batch_size=32, shuffle=False, num_workers=4)
二、内存高效的归一化方案
避免全量数据加载计算统计量,采用分块遍历计算全局均值/方差,再在线应用归一化:
- 分块计算全局统计量:通过内存映射分批次读取数据,累加计算均值和方差,全程仅加载单个batch的数据到内存。
- 在线归一化:在Dataset的
__getitem__中对单个样本应用预计算的均值和方差,或在collate_fn中对batch进行归一化。
示例代码(计算均值和方差):
import numpy as np # 内存映射加载数据集 data_mmap = np.memmap( "your_dataset_path.npy", dtype='float32', mode='r', shape=(64152, 3, 5, 2, 64, 144) ) # 初始化统计量(针对通道维度,即第1维) mean = np.zeros(3, dtype='float64') var = np.zeros(3, dtype='float64') total_samples = 0 batch_size = 1024 # 根据内存容量调整批次大小 # 计算均值 for start in range(0, len(data_mmap), batch_size): end = min(start + batch_size, len(data_mmap)) batch = data_mmap[start:end] # 沿样本、5、2、64、144维度求均值 batch_mean = batch.mean(axis=(0, 2, 3, 4, 5)) mean += batch_mean * (end - start) total_samples += (end - start) mean /= total_samples # 计算方差 for start in range(0, len(data_mmap), batch_size): end = min(start + batch_size, len(data_mmap)) batch = data_mmap[start:end] # 广播均值维度,计算方差 centered = batch - mean.reshape(1, 3, 1, 1, 1, 1) batch_var = (centered ** 2).mean(axis=(0, 2, 3, 4, 5)) var += batch_var * (end - start) var /= total_samples std = np.sqrt(var) # 将统计量传入数据集 full_dataset = MemMapImageDataset("your_dataset_path.npy", mean=mean, std=std)
关键注意事项
- 洗牌仅操作索引,原始数据全程驻留磁盘,彻底避免全量加载的内存峰值。
- 计算统计量时使用
float64类型累加,避免精度损失。 - 若使用多进程DataLoader,确保内存映射文件路径为绝对路径,避免子进程路径问题。
内容的提问来源于stack exchange,提问作者Clayton Malott
相关产品推荐
相关产品推荐

