能否让Python multiprocessing的spawn达到fork的内存效率?
问题:Spawn模式下多进程共享Numpy数组效率优化
我在Linux系统上有一套基于fork的多进程运行代码,利用全局变量共享Numpy数组——由于数组未被修改,依赖Linux的写时复制(COW)机制,不会产生数据拷贝,效率很高。但Python 3.14中Linux平台的multiprocessing将默认使用spawn启动方式,我编写的spawn版本不仅运行更慢,内存占用也更高,推测是Numpy数组被不必要地复制了,希望让spawn版本达到fork版本的效率。
Fork版本代码
from multiprocessing import Pool from time import perf_counter as now import numpy as np def make_func(): n = 20000 np.random.seed(6) start = now() M = np.random.rand(n, n) return lambda x, y: M[x, x] + M[y, y] class ParallelProcessor: def __init__(self): pass def process_task(self, args): """Unpack arguments internally""" index, integer_arg = args print(f(index, integer_arg)) def run_parallel(self, tasks, num_cores=None): """Simplified parallel execution without partial""" num_cores = num_cores task_args = [(idx, val) for idx, val in enumerate(tasks)] start = now() global f f = make_func() print(f"************** {now() - start} seconds to make f") start = now() with Pool(num_cores) as pool: results = pool.map( self.process_task, task_args) print(f"************** {now() - start} seconds to run all jobs") return results if __name__ == "__main__": processor = ParallelProcessor() processor.run_parallel(tasks=[1, 2, 3, 4, 5, 6], num_cores=6)
Spawn版本代码(存在效率问题)
from multiprocessing import Pool, RawArray, set_start_method from time import perf_counter as now import numpy as np def init_worker(shared_array_base, shape): """Initializer function to set up shared memory for each worker""" global M M = np.frombuffer(shared_array_base, dtype=np.float64).reshape(shape) def worker_task(args): """Worker function that reconstructs f using shared memory""" index, integer_arg = args result = M[index, index] + M[integer_arg, integer_arg] print(result) return result class ParallelProcessor: def __init__(self): pass def run_parallel(self, tasks, num_cores=None): """Run tasks in parallel using spawn and shared memory""" set_start_method('spawn', force=True) # Ensure 'spawn' is used n = 20000 shape = (n, n) # Use 'd' for double-precision float (float64) instead of np.float64 shared_array_base = RawArray('d', n * n) M_local = np.frombuffer(shared_array_base, dtype=np.float64).reshape(shape) # Initialize the array in the main process np.random.seed(7) start = now() M_local[:] = np.random.rand(n, n) print(f"************** {now() - start} seconds to make M") # Prepare arguments for worker tasks task_args = [(idx, val) for idx, val in enumerate(tasks)] start = now() with Pool(num_cores, initializer=init_worker, initargs=(shared_array_base, shape)) as pool: results = pool.map(worker_task, task_args) print(f"************** {now() - start} seconds to run all jobs") return results if __name__ == "__main__": processor = ParallelProcessor() processor.run_parallel(tasks=[1, 2, 3, 4, 5, 6], num_cores=6)
解决方案
1. 消除临时数组拷贝
当前代码中M_local[:] = np.random.rand(n, n)会先创建一个独立的临时数组,再将数据拷贝到共享内存,额外增加了一次内存分配和数据复制。直接将随机数生成到共享内存数组中,避免临时数组:
# 替换原数组初始化代码 np.random.seed(7) start = now() np.random.rand(n, n, out=M_local) # 直接写入共享内存,无临时数组开销 print(f"************** {now() - start} seconds to make M")
2. 优化进程初始化逻辑
Spawn模式下每个子进程都会重新加载Python环境,重复调用set_start_method会带来额外开销,将其移到主入口块中:
if __name__ == "__main__": set_start_method('spawn', force=True) # 仅在主进程调用一次 processor = ParallelProcessor() processor.run_parallel(tasks=[1, 2, 3, 4, 5, 6], num_cores=6)
同时移除run_parallel方法中的set_start_method调用。
3. 利用内存映射文件模拟COW效果
对于超大数组,使用np.memmap直接将文件映射到内存,主进程写入后子进程直接读取同一个映射文件,完全避免进程间数据拷贝,效果接近fork的COW:
# 主进程中创建内存映射文件并写入数据 def run_parallel(self, tasks, num_cores=None): n = 20000 shape = (n, n) memmap_path = "/tmp/temp_memmap.dat" # 创建可写的内存映射 M_local = np.memmap(memmap_path, dtype=np.float64, mode='w+', shape=shape) np.random.seed(7) start = now() np.random.rand(n, n, out=M_local) M_local.flush() # 确保数据写入磁盘 print(f"************** {now() - start} seconds to make M") # 子进程初始化函数:只读打开同一个映射文件 def init_worker(path, shape): global M M = np.memmap(path, dtype=np.float64, mode='r', shape=shape) task_args = [(idx, val) for idx, val in enumerate(tasks)] start = now() with Pool(num_cores, initializer=init_worker, initargs=(memmap_path, shape)) as pool: results = pool.map(worker_task, task_args) print(f"************** {now() - start} seconds to run all jobs") # 清理临时文件 import os os.unlink(memmap_path) return results
4. 验证内存共享状态
使用htop或ps aux查看主进程和子进程的内存占用:
- 如果是真正的共享内存,子进程的RSS(常驻内存)不会显著增加,仅VIRT(虚拟内存)会显示相同的数组大小
- 若RSS翻倍,说明代码中存在意外修改数组的操作(触发写时复制),需检查是否有对
M的写入操作
内容的提问来源于stack exchange,提问作者Simd
相关产品推荐
相关产品推荐

