PyMC V5使用Blackjax采样自定义对数似然模型时的JAX转换错误及加速优化求助
PyMC V5使用Blackjax采样自定义对数似然模型时的JAX转换错误及加速优化求助
问题背景
我目前在用PyMC v5实现一个哈密顿蒙特卡洛(HMC)采样的宇宙学模型,为了自定义带梯度的对数似然,我写了一个LogLikeWithGrad的PyTensor Op。现在代码能正常运行,但采样速度极慢——哪怕开了64核并行也没明显改善。
我本来想利用Blackjax的GPU后端来加速采样,于是把原来的采样代码替换成了trace = pm.sample(nuts_sampler="blackjax"),结果直接触发了JAX转换错误,提示无法识别我自定义的Op。想问问有没有大佬遇到过类似问题,或者有其他加速采样的思路?
核心代码片段
自定义对数似然与梯度Op
import pytensor.tensor as pt import scipy.optimize import numpy as np from scipy.optimize import approx_fprime # 带梯度的对数似然Op class LogLikeWithGrad(pt.Op): itypes = [pt.dvector] # 输入为参数向量 otypes = [pt.dscalar] # 输出为单个标量对数似然值 def __init__(self, loglike,): self.likelihood = loglike self.loglike_grad = LogLikeGrad() def perform(self, node, inputs, outputs): (theta,) = inputs logl = self.likelihood(theta,) outputs[0][0] = np.array(logl) def grad(self, inputs, grad_outputs): (theta,) = inputs grads = self.loglike_grad(theta) return [grad_outputs[0] * grads] # 梯度计算Op(用scipy数值梯度) class LogLikeGrad(pt.Op): itypes = [pt.dvector] otypes = [pt.dvector] def __init__(self, ): pass def perform(self, node, inputs, outputs): (theta,) = inputs grads = approx_fprime(theta, applyMCMC, epsilon=1e-8) outputs[0][0] = grads
PyMC模型与采样代码
import pymc as pm import pytensor param_names = ["Omega_m", "Omega_k", "H0", "Psi_0", "dPsi_0_dt", "omega_BD"] logl = LogLikeWithGrad(applyMCMC) model = pm.Model() initial_values = { 'Omega_m': 0.3, 'Omega_k': 1e-3, 'H0' : 67.4, 'Psi_0' : 1.0, 'dPsi_0_dt' : 0.5e-3, 'omega_BD' : 0.5e5, } if __name__ == '__main__': with model: # 定义参数先验分布 for i, name in enumerate(param_names): pm.Uniform(name, lower=lower_boundaries[0][i], upper=upper_boundaries[0][i]) theta = pt.as_tensor_variable([model[param] for param in param_names]) pm.Potential("likelihood", logl(theta)) # 原采样方式(可运行但速度极慢) # niter = 1000 # start = pm.find_MAP() # step = pm.NUTS() # trace = pm.sample(draws=niter, step=step, tune=500, cores=64, init="jitter+adapt_diag", progressbar=True) # 尝试用Blackjax加速,触发报错的代码 trace = pm.sample(nuts_sampler="blackjax")
补充说明
我的applyMCMC函数逻辑比较复杂:需要调用Hi-CLASS计算宇宙学功率谱,结合Planck似然、哈勃常数观测数据、距离模数数据计算总对数似然,中间还包含RK4求解微分方程的过程——这部分应该是采样慢的核心原因。
报错信息
return fgraph_to_python( File "/opt/intel/oneapi/intelpython/python3.9/lib/python3.9/site-packages/pytensor/link/utils.py", line 734, in fgraph_to_python compiled_func = op_conversion_fn( File "/opt/intel/oneapi/intelpython/python3.9/lib/python3.9/functools.py", line 888, in wrapper return dispatch(args[0].__class__)(*args, **kw) File "/opt/intel/oneapi/intelpython/python3.9/lib/python3.9/site-packages/pytensor/link/jax/dispatch/basic.py", line 41, in jax_funcify raise NotImplementedError(f"No JAX conversion for the given `Op`: {op}") NotImplementedError: No JAX conversion for the given `Op`: LogLikeWithGrad
求助方向
- 有没有办法让Blackjax识别我自定义的PyTensor Op?或者有没有其他方式在PyMC v5.10里结合Blackjax做GPU加速?
- 除了换采样器,针对我这种包含大量数值计算(RK4、外部宇宙学库调用)的似然函数,还有什么优化方案能提升采样速度?
备注:内容来源于stack exchange,提问作者foutou_10
相关产品推荐
相关产品推荐

