体素Fisher精确检验的多进程优化实现问题咨询
问题描述
我要对尺寸为(240,240,155)的两幅3D图像的每个体素执行Fisher精确检验,只会基础的多进程用法,不知道怎么把多进程用到这个场景里。
我尝试的原始代码:
result_array = np.zeros_like(imgRHUH, dtype=float) for (x_RHUH,y_RHUH,z_RHUH), valueRHUH in np.ndenumerate(imgRHUH): for (x_BTP,y_BTP,z_BTP), valueBTP in np.ndenumerate(imgBTP): a = int(valueRHUH) b = int(valueBTP) c = 39 - a d = 20 - b n = 2 p = (math.factorial(a + b)*math.factorial(c + d)*math.factorial(a + c)*math.factorial(b + d)) / (math.factorial(a)*math.factorial(b)*math.factorial(c)*math.factorial(d)*math.factorial(n)) result_array[x_RHUH, y_RHUH, z_RHUH] = p
我曾尝试按如下方式分块处理第一幅图像,但性能提升不明显,还担心合并时出现顺序错误:
chunks = [] for chunk in np.array_split(list(np.ndenumerate(imgRHUH)), multiprocessing.cpu_count()): chunks.append((chunk, imgRHUH, imgBTP, result_array))
解决方案
先修复核心计算逻辑
你的原始代码存在致命冗余:每个体素的计算只需要两幅图像相同位置的体素值,但你嵌套遍历了两幅图的所有体素,计算量直接从O(N)飙升到O(N²)(N是体素总数),这才是性能差的核心原因。先把逻辑改成单循环:
import numpy as np import math def compute_fisher(v_rhuh, v_btp): a = int(v_rhuh) b = int(v_btp) c = 39 - a d = 20 - b n = 2 p = (math.factorial(a + b)*math.factorial(c + d)*math.factorial(a + c)*math.factorial(b + d)) / \ (math.factorial(a)*math.factorial(b)*math.factorial(c)*math.factorial(d)*math.factorial(n)) return p # 正确的单循环版本 result_array = np.zeros_like(imgRHUH, dtype=float) for idx in np.ndindex(imgRHUH.shape): result_array[idx] = compute_fisher(imgRHUH[idx], imgBTP[idx])
多进程优化方案
针对3D图像,按轴拆分切片(比如z轴)是最合理的方式,既不会打乱索引顺序,又能保证每个进程处理连续内存块,提升效率。
方案1:基于multiprocessing.Pool的切片并行
import numpy as np import math import multiprocessing def process_slice(slice_rhuh, slice_btp): slice_result = np.zeros_like(slice_rhuh, dtype=float) for idx in np.ndindex(slice_rhuh.shape): a = int(slice_rhuh[idx]) b = int(slice_btp[idx]) c = 39 - a d = 20 - b n = 2 p = (math.factorial(a + b)*math.factorial(c + d)*math.factorial(a + c)*math.factorial(b + d)) / \ (math.factorial(a)*math.factorial(b)*math.factorial(c)*math.factorial(d)*math.factorial(n)) slice_result[idx] = p return slice_result if __name__ == "__main__": # 按CPU核心数拆分z轴切片 num_workers = multiprocessing.cpu_count() rhuh_slices = np.array_split(imgRHUH, num_workers, axis=2) btp_slices = np.array_split(imgBTP, num_workers, axis=2) # 并行计算 with multiprocessing.Pool(num_workers) as pool: slice_results = pool.starmap(process_slice, zip(rhuh_slices, btp_slices)) # 合并切片得到最终结果 result_array = np.concatenate(slice_results, axis=2)
方案2:共享内存优化(超大数组适用)
如果图像体积过大,传递切片会产生内存拷贝开销,可以用共享内存避免重复拷贝:
import numpy as np import math import multiprocessing from multiprocessing import shared_memory def process_shared(shm_rhuh_name, shm_btp_name, shm_result_name, shape, z_range): # 连接共享内存 shm_rhuh = shared_memory.SharedMemory(name=shm_rhuh_name) shm_btp = shared_memory.SharedMemory(name=shm_btp_name) shm_result = shared_memory.SharedMemory(name=shm_result_name) # 关联numpy数组 imgRHUH = np.ndarray(shape, dtype=imgRHUH.dtype, buffer=shm_rhuh.buf) imgBTP = np.ndarray(shape, dtype=imgBTP.dtype, buffer=shm_btp.buf) result_array = np.ndarray(shape, dtype=np.float64, buffer=shm_result.buf) # 处理指定z轴范围 z_start, z_end = z_range for z in range(z_start, z_end): for y in range(shape[1]): for x in range(shape[0]): a = int(imgRHUH[x,y,z]) b = int(imgBTP[x,y,z]) c = 39 - a d = 20 - b n = 2 p = (math.factorial(a + b)*math.factorial(c + d)*math.factorial(a + c)*math.factorial(b + d)) / \ (math.factorial(a)*math.factorial(b)*math.factorial(c)*math.factorial(d)*math.factorial(n)) result_array[x,y,z] = p # 关闭共享内存连接 shm_rhuh.close() shm_btp.close() shm_result.close() if __name__ == "__main__": shape = imgRHUH.shape dtype_result = np.float64 # 创建共享内存并写入数据 shm_rhuh = shared_memory.SharedMemory(create=True, size=imgRHUH.nbytes) shm_btp = shared_memory.SharedMemory(create=True, size=imgBTP.nbytes) shm_result = shared_memory.SharedMemory(create=True, size=np.zeros(shape, dtype=dtype_result).nbytes) imgRHUH_shared = np.ndarray(shape, dtype=imgRHUH.dtype, buffer=shm_rhuh.buf) imgRHUH_shared[:] = imgRHUH[:] imgBTP_shared = np.ndarray(shape, dtype=imgBTP.dtype, buffer=shm_btp.buf) imgBTP_shared[:] = imgBTP[:] # 拆分z轴任务 num_workers = multiprocessing.cpu_count() slice_size = shape[2] // num_workers z_ranges = [] for i in range(num_workers): start = i * slice_size end = start + slice_size if i != num_workers-1 else shape[2] z_ranges.append((start, end)) # 启动进程 processes = [] for z_range in z_ranges: p = multiprocessing.Process(target=process_shared, args=(shm_rhuh.name, shm_btp.name, shm_result.name, shape, z_range)) processes.append(p) p.start() # 等待进程完成 for p in processes: p.join() # 获取结果 result_array = np.copy(np.ndarray(shape, dtype=dtype_result, buffer=shm_result.buf)) # 清理共享内存 shm_rhuh.close() shm_rhuh.unlink() shm_btp.close() shm_btp.unlink() shm_result.close() shm_result.unlink()
额外建议
- 验证Fisher公式正确性:建议用
scipy.stats.fisher_exact对比结果,确保手动实现的公式逻辑正确,比如:from scipy.stats import fisher_exact oddsratio, p_value = fisher_exact([[a, b], [c, d]]) - 尝试向量化运算:可以把计算逻辑改成numpy数组操作,完全替代循环,进一步提升单进程性能。
内容的提问来源于stack exchange,提问作者Gaspar Alves Gonçalves
相关产品推荐
相关产品推荐

