如何高效读取多个NPY文件以避免内存溢出?
内存优化方案与代码修改建议
问题根源
你用allow_pickle=True加载的是Python字典组成的numpy数组,每个字典都是独立的Python对象,内存开销远大于磁盘文件大小(比如2.2GB的磁盘文件,加载后内存可能涨到10GB以上);再加上np.concatenate每次都会创建新数组,瞬间占用双倍内存,直接触发系统OOM杀进程。
核心解决方法:转用Numpy结构化数组
把原来的字典数组转成numpy原生的结构化数组,消除Python对象的内存开销,内存占用会和磁盘大小接近。
第一步:转换现有文件格式
先写个脚本,把现有的pickle格式NPY转成结构化数组:
import numpy as np import os DATA_DIR = DATA_MODEL_NON_UNIFORM_1_MLN filenames = [f for f in os.listdir(DATA_DIR) if f.endswith('.npy')] for fname in filenames: file_path = os.path.join(DATA_DIR, fname) # 加载原pickle格式数据 raw_data = np.load(file_path, allow_pickle=True) # 从第一个字典生成结构化数组的dtype sample = raw_data[0] dtype_def = [] for key, val in sample.items(): arr = np.array(val) # 每个字段的格式:(字段名, 数据类型, 数组形状) dtype_def.append((key, arr.dtype, arr.shape)) # 构建结构化数组 structured_data = np.empty(len(raw_data), dtype=dtype_def) for idx, item in enumerate(raw_data): for key in item: structured_data[key][idx] = item[key] # 保存为新的NPY文件(加后缀区分原文件) new_path = os.path.join(DATA_DIR, f"structured_{fname}") np.save(new_path, structured_data)
第二步:优化加载代码
转格式后,用预分配数组的方式加载合并,避免concatenate的双倍内存开销:
import numpy as np import os def load_all_data(data_dir): # 筛选转好的结构化文件 filenames = [f for f in os.listdir(data_dir) if f.startswith('structured_') and f.endswith('.npy')] if not filenames: raise ValueError("未找到结构化格式的NPY文件") # 获取单文件的dtype和长度 first_path = os.path.join(data_dir, filenames[0]) first_data = np.load(first_path, mmap_mode='r') total_samples = len(first_data) * len(filenames) dtype = first_data.dtype # 预分配总数组 all_data = np.empty(total_samples, dtype=dtype) # 写入数据 current_idx = 0 for fname in filenames: file_path = os.path.join(data_dir, fname) # 用mmap加载,减少内存峰值 chunk = np.load(file_path, mmap_mode='r') chunk_len = len(chunk) # 直接切片赋值,无额外内存开销 all_data[current_idx:current_idx+chunk_len] = chunk[:] current_idx += chunk_len del chunk return all_data
可选优化方案
- 拆分字段存储:如果不需要同时使用所有字段,可以把每个字段单独存为一个NPY文件(比如
large_list1.npy、small_list1.npy),加载时只读取需要的字段,进一步降低内存占用。 - 改用HDF5格式:用
h5py库把数据存为HDF5文件,支持分块存储、按需加载,甚至可以不把全部数据读入内存,直接在磁盘上操作,适合超大规模数据。示例代码:import h5py import numpy as np # 写入HDF5 with h5py.File('data.h5', 'w') as f: # 假设已经有dtype_def和total_samples变量 for key in dtype_def: field_name = key[0] # 创建分块存储的dataset f.create_dataset(field_name, shape=(total_samples,) + key[2], dtype=key[1], chunks=(1000,) + key[2]) # 逐个文件写入数据 current_idx = 0 for fname in filenames: chunk = np.load(os.path.join(DATA_DIR, fname), mmap_mode='r') chunk_len = len(chunk) f[field_name][current_idx:current_idx+chunk_len] = chunk[field_name][:] current_idx += chunk_len # 读取HDF5(按需加载) with h5py.File('data.h5', 'r') as f: # 只读取前1000条的large_list1字段 partial_data = f['large_list1'][:1000]
临时应急方案(不转格式)
如果暂时没时间转格式,不要用np.concatenate,而是用列表extend的方式合并,但这种方法内存占用还是很高,仅作临时用:
import numpy as np import os def load_all_raw(data_dir): filenames = [f for f in os.listdir(data_dir) if f.endswith('.npy')] all_data = [] for fname in filenames: file_path = os.path.join(data_dir, fname) # 用mmap加载,减少瞬间内存占用 chunk = np.load(file_path, allow_pickle=True, mmap_mode='r') all_data.extend(chunk.tolist()) del chunk # 最后转成numpy数组(如果必须的话) return np.array(all_data, dtype=object)
内容的提问来源于stack exchange,提问作者mkow93
相关产品推荐
相关产品推荐

