PyTorch如何加载多份大容量pickle数据集训练避免内存问题
大体积分块Pickle数据集训练防内存溢出方案
核心原则是永远不要全量加载所有数据到内存,用按需读取+轻量缓存的逻辑,让内存里同一时间只保留当前训练需要用到的最少数据,哪怕总数据量远大于内存容量也能正常跑训练。
方案1:自定义映射式数据集(零额外预处理,开箱即用)
这是改造成本最低的方案,不需要改动原始数据,适配PyTorch、TensorFlow等所有主流训练框架。
具体实现步骤
- 第一步:预生成轻量样本索引
不要在训练初始化阶段直接加载所有pickle文件内容,先遍历所有pickle文件,仅读取每个文件内x、y列表的长度,为每一条(x,y)样本生成全局索引,记录「样本所属文件路径、样本在文件内列表的下标」,把这个索引表单独存成几MB大小的json/小pickle文件,训练启动时只加载这个索引表,几乎不占内存。
索引生成参考代码:
import os import pickle import json DATASET_ROOT = "./dataset_folder" INDEX_SAVE_PATH = "./dataset_index.json" index_records = [] # 遍历目录下所有pickle文件 for file_name in os.listdir(DATASET_ROOT): if not file_name.endswith(".pickle"): continue file_path = os.path.join(DATASET_ROOT, file_name) # 临时打开文件仅读取列表长度,读完立刻释放文件句柄和内容 with open(file_path, "rb") as f: x_list, y_list = pickle.load(f) assert len(x_list) == len(y_list), f"文件{file_name}中x、y列表长度不匹配" sample_num = len(x_list) # 记录当前文件内所有样本的位置 for inner_idx in range(sample_num): index_records.append({ "file_path": file_path, "inner_idx": inner_idx }) # 保存索引表 with open(INDEX_SAVE_PATH, "w") as f: json.dump(index_records, f)
- 第二步:实现带单文件缓存的数据集类
数据集初始化阶段只加载刚才生成的轻量索引,维护两个缓存变量记录当前已经加载到内存的pickle文件路径、对应的x和y列表。当需要读取某一条样本时:- 如果样本所属文件就是当前缓存的文件,直接从缓存里按下标取数据
- 如果不属于当前缓存文件,立刻清空旧缓存的内容,打开目标pickle文件加载内容到缓存,再取对应样本
这个逻辑下,内存里永远最多只存1个pickle文件的内容,哪怕总数据量上百G也不会撑爆内存。
PyTorch框架下的数据集参考实现:
import json import pickle from torch.utils.data import Dataset class ChunkedPickleDataset(Dataset): def __init__(self, index_path=INDEX_SAVE_PATH): # 仅加载索引表,内存占用可忽略 with open(index_path, "r") as f: self.index = json.load(f) # 初始化缓存为空 self.cached_path = None self.cached_x = None self.cached_y = None def __len__(self): return len(self.index) def __getitem__(self, idx): sample_meta = self.index[idx] target_path = sample_meta["file_path"] inner_idx = sample_meta["inner_idx"] # 缓存命中直接返回 if target_path == self.cached_path: return self.cached_x[inner_idx], self.cached_y[inner_idx] # 缓存未命中,先清空旧缓存 self.cached_path = None self.cached_x = None self.cached_y = None # 加载目标文件到缓存 with open(target_path, "rb") as f: x_list, y_list = pickle.load(f) self.cached_path = target_path self.cached_x = x_list self.cached_y = y_list return x_list[inner_idx], y_list[inner_idx]
- 第三步:配合数据加载器的参数调优
用框架自带的DataLoader/DataSet加载器时,num_workers建议设为2-4(根据CPU核心数调整),不要开太大——每个worker进程会独立维护自己的文件缓存,开太多会导致多进程同时加载多个文件,额外占用内存。普通消费级显卡训练时开pin_memory=True即可,不需要额外做内存预占。
方案2:格式转换(适合长期反复训练的场景)
如果这套数据集需要多次迭代训练,可以花一次预处理时间,把所有pickle里的样本转成天生支持流式随机读取的格式,比如LMDB、WebDataset或者numpy内存映射格式,后续训练时不需要做pickle反序列化,加载速度能提升30%以上,内存占用更低。
注意不要把所有数据合并成单个大pickle文件,大pickle必须全量加载到内存才能读取,反而更容易触发OOM。
避坑要点
- 绝对不要在训练启动时写循环把所有pickle文件全量加载到一个大列表里:15GB的pickle反序列化后,Python对象的实际内存占用会达到原始文件大小的2-3倍,32G内存的机器会直接OOM。
- 如果单个pickle文件本身大小就超过机器可用内存,需要先把单个大pickle拆分成更小的分块,再用上面的缓存方案加载。
- 训练过程中的数据增强结果不要缓存到内存,处理完直接送入模型即可。
内容的提问来源于stack exchange,提问作者Bshr Eldebuch
相关产品推荐
相关产品推荐

