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

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列表。当需要读取某一条样本时:
    1. 如果样本所属文件就是当前缓存的文件,直接从缓存里按下标取数据
    2. 如果不属于当前缓存文件,立刻清空旧缓存的内容,打开目标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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 08:33:24