You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.20 13:03:12