如何并行化scipy.integrate.quad加速后验概率的卡方计算?
提速卡方值计算(用于后验概率估计)
问题背景
计算参数θ的后验概率时,似然函数的卡方值计算速度成为瓶颈——每个z_sn样本都需要单独执行一次积分运算。当前使用列表推导式处理样本循环,不确定是否为最优方案,了解Cython但未实践,使用emcee进行后验概率计算。
原始代码
import numpy as np import scipy.constants as cte from scipy.integrate import quad def luminosity_integrand(z, omgM): Ez = np.sqrt((1 - omgM) + omgM * np.power(1 + z, 3)) return 1. / Ez def luminosity_distance(z, h, omgM): integral, _ = quad(luminosity_integrand, 0, z, epsrel=1e-8, args=(omgM)) return (cte.c / 10. ** 5) / h * (1 + z) * integral def distance_modulus(z, h, omgM): return 5. * np.log10(luminosity_distance(z, h, omgM)) + 25. def chisq_sn(h, omgM): m_model = np.array([distance_modulus(z, h, omgM) for z in z_sn]) diffs = m_obs-m_model maha_distances = np.dot(np.dot(diffs, inv_cov_plus), diffs) # mahalanobis distance return maha_distances
优化方案
1. 改用向量化积分(最直接有效)
scipy.integrate.quad是单值积分函数,循环调用会带来大量Python层开销。使用scipy.integrate.quad_vec(Scipy 1.9+版本支持)可一次性处理所有z_sn样本,利用向量化运算大幅提速:
修改核心函数:
from scipy.integrate import quad_vec def luminosity_distance(z_arr, h, omgM): # z_arr是z_sn的numpy数组,一次性计算所有积分 integral, _ = quad_vec(luminosity_integrand, 0, z_arr, epsrel=1e-8, args=(omgM,)) return (cte.c / 10.**5) / h * (1 + z_arr) * integral def distance_modulus(z_arr, h, omgM): return 5. * np.log10(luminosity_distance(z_arr, h, omgM)) + 25. def chisq_sn(h, omgM): # 直接传入z_sn数组,无需列表推导式 m_model = distance_modulus(z_sn, h, omgM) diffs = m_obs - m_model maha_distances = np.dot(np.dot(diffs, inv_cov_plus), diffs) return maha_distances
2. 预计算积分表(适合参数范围有限的场景)
如果emcee采样时omgM的取值范围有限,可以提前在所有可能的omgM值下,预计算所有z_sn对应的积分结果,存储为二维数组。后续计算时直接查表,彻底避免重复积分:
# 预计算示例 omgM_grid = np.linspace(0.1, 0.5, 100) # 根据实际采样范围调整 integral_cache = np.zeros((len(omgM_grid), len(z_sn))) for i, omgM in enumerate(omgM_grid): integral_cache[i] = quad_vec(luminosity_integrand, 0, z_sn, epsrel=1e-8, args=(omgM,))[0] # 计算时通过插值获取对应omgM的积分值 def get_integral(omgM): return np.interp(omgM, omgM_grid, integral_cache, axis=0)
3. 放宽积分精度要求
当前设置的epsrel=1e-8精度较高,可尝试放宽到1e-6或1e-7,积分速度会显著提升。需验证:精度降低后,卡方值的变化是否会影响后验概率的最终结果。
4. Cython加速(进阶优化)
若上述方法仍达不到速度要求,可将核心计算逻辑用Cython重写,减少Python解释器的开销。示例Cython代码结构:
# cython: language_level=3 import numpy as np cimport numpy as np from scipy.integrate cimport quad cdef double luminosity_integrand_c(double z, double omgM): cdef double Ez = sqrt((1 - omgM) + omgM * pow(1 + z, 3)) return 1.0 / Ez def luminosity_distance_c(double z, double h, double omgM): cdef double integral, err integral, err = quad(luminosity_integrand_c, 0, z, epsrel=1e-8, args=(omgM,)) # 直接写入光速数值,避免调用scipy.constants的开销 return (299792.458 / 100000.0) / h * (1 + z) * integral
编译后替换原luminosity_distance函数即可,可配合Cython的编译优化选项进一步提速。
5. 利用emcee的多进程并行
emcee支持多进程并行计算,通过multiprocessing.Pool将不同walker的计算分配到多核CPU上,充分利用硬件资源:
from multiprocessing import Pool import emcee def log_probability(params): h, omgM = params # 这里调用chisq_sn,返回对数后验概率 return -0.5 * chisq_sn(h, omgM) # 初始化并行池 with Pool() as pool: sampler = emcee.EnsembleSampler(nwalkers=32, ndim=2, log_prob_fn=log_probability, pool=pool) sampler.run_mcmc(initial_state, nsteps=1000)
内容的提问来源于stack exchange,提问作者Alucard
相关产品推荐
相关产品推荐

