如何用生成器迭代读取大体积object类型numpy保存文件?
这个问题确实戳中了numpy object dtype文件的一个痛点——因为object类型的存储机制和普通数值类型完全不同,mmap_mode确实没法直接用。不过好在有几个可行的方案能实现流式/迭代读取,不用把整个文件塞进内存,我来给你拆解一下:
方案1:手动解析.npy文件格式,逐个读取元素
numpy的.npy文件是有明确结构的,哪怕是object dtype的数组,我们可以手动解析头部,然后逐个读取并反序列化每个元素:
- 读取文件头部:用numpy内置的工具函数先读取.npy的魔法数和头部信息,拿到数组的shape、dtype等关键信息;
- 逐个读取元素:object类型的数组在.npy文件中,每个元素会先存储一个8字节的长度值(标记该元素pickle序列化后的字节数),接着是对应的序列化字节。我们可以循环读取这些内容,反序列化后返回单个元素。
具体代码实现:
import numpy as np import pickle def read_npy_object_generator(file_path): with open(file_path, 'rb') as f: # 读取魔法数,确认是.npy文件 magic = np.lib.format.read_magic(f) # 读取头部信息,适配不同版本的.npy格式 if magic == (1, 0): header = np.lib.format.read_array_header_1_0(f) elif magic == (2, 0): header = np.lib.format.read_array_header_2_0(f) else: raise ValueError("Unsupported .npy file version") dtype = np.dtype(header['descr']) shape = header['shape'] # 确认是object dtype if dtype.kind != 'O': raise ValueError("This generator only works for object dtype arrays") # 计算总元素数 total_elements = np.prod(shape) for _ in range(total_elements): # 读取元素的字节长度(numpy用uint64存储) length_bytes = f.read(8) if not length_bytes: break # 意外到达文件末尾,提前终止 element_length = np.frombuffer(length_bytes, dtype=np.uint64)[0] # 读取对应长度的序列化字节 element_bytes = f.read(element_length) # 反序列化得到原始元素 element = pickle.loads(element_bytes) yield element
使用这个生成器时,你可以用for elem in read_npy_object_generator('your_file.npy'):来逐个处理元素,每次内存里只会加载当前元素,完全不用加载整个大数组。
方案2:保存时提前拆分,用分块存储格式
如果还没保存这个大数组,或者可以重新保存的话,换用支持分块/流式读取的存储格式会更省心:
- HDF5(h5py):可以把每个变长数组作为单独的dataset,或者用一个group来统一管理。读取时逐个遍历dataset即可:
import h5py # 保存大数组的示例代码 with h5py.File('large_object_array.h5', 'w') as f: elem_group = f.create_group('elements') for idx, sub_arr in enumerate(your_large_object_array): elem_group.create_dataset(f'elem_{idx}', data=sub_arr) # 读取时的生成器 def h5_object_generator(file_path): with h5py.File(file_path, 'r') as f: elem_group = f['elements'] # 按索引排序确保顺序正确 for key in sorted(elem_group.keys(), key=lambda x: int(x.split('_')[1])): yield elem_group[key][()]
- Zarr:和HDF5类似,但更适合云存储和并行场景,同样支持分块存储,读取时可以按需加载单个元素,不需要一次性加载全部数据。
方案3:用Dask处理超大规模数组
如果你的变长数组有一定规律性,或者可以转换为Dask支持的结构,Dask可以帮你实现延迟加载和流式处理:
import dask.bag as db # 基于方案1的生成器创建Dask Bag element_bag = db.from_sequence(read_npy_object_generator('your_file.npy'), npartitions=10) # 可以对Bag进行map、filter等操作,所有操作都是延迟执行的,不会一次性加载所有元素 example_result = element_bag.map(lambda arr: arr.mean()).compute()
不过要注意,Dask对object dtype的支持不如数值类型完善,要是元素是复杂嵌套结构,可能需要额外做适配处理。
最后补充一句:如果你的object数组里的单个元素本身就是很大的numpy数组,那单个元素还是会占用不少内存,但至少不会把整个大数组一次性加载到内存里,能有效缓解内存压力。
内容的提问来源于stack exchange,提问作者David Parks
相关产品推荐
相关产品推荐

