将大型numpy/TensorFlow/PyTorch数组缓存至磁盘以省内存保性能
多进程大数组聚合的磁盘缓存优化方案
核心思路:绕过通用序列化,直接操作二进制数据
DiskCache的通用序列化(pickle)对超大数组效率极低,因为要把整个对象序列化/反序列化,不如直接把数组的二进制数据写入磁盘,同时用异步IO提前缓冲。
一、分块写入+内存映射(memmap)
针对不同框架的数组,直接导出二进制数据,用内存映射文件实现磁盘与内存的直接映射,避免全量加载:
- NumPy数组:
# 预先创建空的内存映射文件(主进程执行) import numpy as np total_shape = (10000, 10000) dtype = np.float32 file_size = np.prod(total_shape) * dtype.itemsize with open('large_array.bin', 'wb') as f: f.seek(file_size - 1) f.write(b'\0') # 多进程写入各自负责的块 memmap = np.memmap('large_array.bin', dtype=dtype, mode='r+', shape=total_shape) # 假设当前进程负责第i块,计算start/end索引 start_idx = i * 2500 end_idx = (i+1) * 2500 memmap[start_idx:end_idx] = local_numpy_array # local_numpy_array为当前进程的数组块 memmap.flush() # 刷入磁盘 # 聚合进程读取,无需全量加载 agg_memmap = np.memmap('large_array.bin', dtype=dtype, mode='r', shape=total_shape) # 按需读取块或直接处理 block = agg_memmap[0:2500] - PyTorch/TensorFlow数组:先转成NumPy数组再用memmap,或直接导出二进制:
PyTorch示例:# 进程内写入 local_tensor = torch.randn(1000, 1000) np_memmap = np.memmap('torch_block.bin', dtype=np.float32, mode='w+', shape=local_tensor.shape) np_memmap[:] = local_tensor.numpy() np_memmap.flush() # 聚合进程读取 np_memmap = np.memmap('torch_block.bin', dtype=np.float32, mode='r', shape=(1000,1000)) agg_tensor = torch.from_numpy(np_memmap)
二、异步IO预缓冲
用线程池在聚合进程中提前异步读取磁盘块,避免IO等待拖慢性能:
from concurrent.futures import ThreadPoolExecutor import numpy as np def read_block(file_path, total_shape, dtype, start, end): memmap = np.memmap(file_path, dtype=dtype, mode='r', shape=total_shape) return memmap[start:end].copy() # 提前缓存到内存 # 聚合进程启动线程池预读 executor = ThreadPoolExecutor(max_workers=4) future_blocks = [] for i in range(4): start = i * 2500 end = (i+1) * 2500 future = executor.submit(read_block, 'large_array.bin', (10000,10000), np.float32, start, end) future_blocks.append(future) # 处理时直接获取预加载好的块 for future in future_blocks: block = future.result() # 执行聚合逻辑(如拼接、求和)
三、专用大型数组存储库替代方案
如果不想手动实现memmap逻辑,可使用专门处理大数组的磁盘存储库:
- zarr:支持分块存储、压缩,多进程安全写入:
import zarr # 多进程写入 store = zarr.DirectoryStore('zarr_storage') z_arr = zarr.open(store, shape=(10000,10000), dtype='f4', mode='w') z_arr[start_idx:end_idx] = local_array # 聚合进程读取 z_arr = zarr.open('zarr_storage', mode='r') block = z_arr[0:2500] # 按需读取 - h5py:基于HDF5格式,支持多进程并行写入(需MPI环境):
import h5py # 多进程写入(需配置MPI) with h5py.File('large_array.h5', 'w', driver='mpio', comm=MPI.COMM_WORLD) as f: dset = f.create_dataset('data', shape=(10000,10000), dtype='f4') dset[start_idx:end_idx] = local_array # 聚合进程读取 with h5py.File('large_array.h5', 'r') as f: block = f['data'][0:2500]
关键注意事项
- 避免用pickle序列化超大数组:通用序列化会把整个数组加载到内存再写入,完全失去磁盘缓存的意义。
- 多进程写入时保证块不重叠:提前规划每个进程负责的数组区间,避免写入冲突。
- 优先使用SSD存储:SSD的随机读写性能接近内存,能大幅降低IO等待时间。
内容的提问来源于stack exchange,提问作者Daniel Wiczew
相关产品推荐
相关产品推荐

