使用multiprocessing加速HDF5数据lmfit高斯拟合遇文件读取错误求助
问题分析与解决方案
问题根源
- Windows多进程启动机制冲突:Windows下
multiprocessing默认用spawn模式,子进程会重新执行整个脚本。你的文件读取逻辑在if __name__ == '__main__'之外,导致每个子进程都尝试重新打开HDF5文件,引发并发访问冲突;加上K盘是网络驱动器,子进程可能无法正常访问该路径。 - 代码冗余与效率问题:
fitting_curve内重复创建gModel,浪费初始化资源;params创建重复:gModel.make_params已生成参数,无需再用params.add重复添加;- 多进程嵌套循环过多,
pool.map使用方式未充分利用批量处理能力; sigma == 'nan'的判断逻辑错误,无法识别数值类型的NaN。
解决方案
方案1:全内存加载+规范多进程流程
适合内存足够容纳4.3GB数据的场景,确保数据完全加载到内存后批量处理:
import numpy as np import h5py import hdf5plugin from lmfit import Model from functools import partial import multiprocessing as mp # 全局定义高斯模型,避免重复创建 gModel = Model( lambda x, amp, x0, wid: (amp / np.sqrt(2*np.pi*wid))*np.exp(-(x-x0)**2/(2*wid**2)), independent_vars=['x'], param_names=['amp','x0','wid'], nan_policy='omit' ) def parameters(x, arr): sum_arr = arr.sum() if sum_arr == 0: return (0, 0, 0) mean = (x * arr).sum() / sum_arr sigma = np.sqrt((arr * (x - mean)**2).sum() / sum_arr) amplitude = arr.max() return (mean, sigma, amplitude) def fitting_curve(x, arr): mean, sigma, amplitude = parameters(x, arr) upper_lim = 10**7 lower_lim = 10**1 if amplitude == 0 or np.isnan(sigma): return 0 params = gModel.make_params(amp=amplitude, x0=mean, wid=sigma) params['wid'].min = 0.0 result = gModel.fit(arr, x=x, params=params) amp_fitted = result.best_values['amp'] return amp_fitted if lower_lim < amp_fitted < upper_lim else 0 if __name__ == '__main__': fname = 'K:/11014463/processed/phantom_offline/waxs/merged_data.nxs' # 仅主进程读取数据,确保完全加载到内存 with h5py.File(fname, mode='r', driver='core') as a: data = np.array(a['processed']['result']['data'][()]) q_vector = np.array(a['processed']['result']['r_binning'][()]) n_cores = mp.cpu_count() print(f"使用 {n_cores-1} 个核心") nr_sec = 4 # 扁平化所有待拟合的一维数组任务 tasks = [data[k,m,o,nr_sec,:] for k in range(data.shape[0]) for m in range(data.shape[1]) for o in range(data.shape[2])] func_p = partial(fitting_curve, q_vector) # 批量处理任务 with mp.Pool(n_cores-1) as pool: results = pool.map(func_p, tasks) # 将结果重塑为3D数组 arr_out = np.array(results).reshape(data.shape[0], data.shape[1], data.shape[2]) print(f"拟合完成,结果数组形状:{arr_out.shape}")
方案2:子进程独立读取HDF5(低内存场景)
如果内存不足以加载全部数据,让每个子进程独立读取所需切片,避免大数组进程间传递:
import numpy as np import h5py import hdf5plugin from lmfit import Model import multiprocessing as mp gModel = Model( lambda x, amp, x0, wid: (amp / np.sqrt(2*np.pi*wid))*np.exp(-(x-x0)**2/(2*wid**2)), independent_vars=['x'], param_names=['amp','x0','wid'], nan_policy='omit' ) def parameters(x, arr): sum_arr = arr.sum() if sum_arr == 0: return (0, 0, 0) mean = (x * arr).sum() / sum_arr sigma = np.sqrt((arr * (x - mean)**2).sum() / sum_arr) amplitude = arr.max() return (mean, sigma, amplitude) def fitting_task(args): fname, q_vector, k, m, o, nr_sec = args # 子进程独立打开文件读取目标切片 with h5py.File(fname, mode='r', driver='core') as a: arr = np.array(a['processed']['result']['data'][k,m,o,nr_sec,:]) mean, sigma, amplitude = parameters(q_vector, arr) upper_lim = 10**7 lower_lim = 10**1 if amplitude == 0 or np.isnan(sigma): return 0 params = gModel.make_params(amp=amplitude, x0=mean, wid=sigma) params['wid'].min = 0.0 result = gModel.fit(arr, x=q_vector, params=params) amp_fitted = result.best_values['amp'] return amp_fitted if lower_lim < amp_fitted < upper_lim else 0 if __name__ == '__main__': fname = 'K:/11014463/processed/phantom_offline/waxs/merged_data.nxs' nr_sec = 4 # 主进程仅读取q_vector和数据维度 with h5py.File(fname, mode='r', driver='core') as a: q_vector = np.array(a['processed']['result']['r_binning'][()]) data_shape = a['processed']['result']['data'].shape n_cores = mp.cpu_count() print(f"使用 {n_cores-1} 个核心") # 生成所有任务参数 tasks = [(fname, q_vector, k, m, o, nr_sec) for k in range(data_shape[0]) for m in range(data_shape[1]) for o in range(data_shape[2])] with mp.Pool(n_cores-1) as pool: results = pool.map(fitting_task, tasks) arr_out = np.array(results).reshape(data_shape[0], data_shape[1], data_shape[2]) print(f"拟合完成,结果数组形状:{arr_out.shape}")
关键优化说明
- 脚本入口规范:所有文件读取、多进程启动逻辑放入
if __name__ == '__main__',避免Windows子进程重复执行初始化代码; - 模型全局化:
gModel全局定义,每个子进程仅初始化一次; - 参数逻辑简化:移除冗余的
params.add调用; - 批量任务处理:扁平化所有拟合任务,用
pool.map批量提交,减少循环开销; - NaN判断修正:用
np.isnan(sigma)替代错误的字符串判断; - 内存适配:方案2适合低内存场景,子进程按需读取数据切片。
内容的提问来源于stack exchange,提问作者alccdesy
相关产品推荐
相关产品推荐

