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

基于Turing.jl的复杂多项式MCMC(NUTS)高效推断优化咨询

复杂多项式模型的Turing.jl NUTS采样优化方案

我在模型中使用多个复杂多项式方程进行评估,希望借助Turing.jl实现高效的MCMC(NUTS)推断。目前用的是加噪声的模拟数据,后续会替换成实验数据。当前案例代码能正常运行,但多项式方程更复杂时采样速度会变得极慢。试过用@generated函数优化,效果不好,求这类场景的最佳实践。

优化实践方案

1. 放弃@generated函数,改用普通函数+向量化计算

@generated函数适合编译期类型固定的场景,你的模型都是数值计算,普通函数配合向量化反而能减少编译开销,提升计算效率。将所有生成式函数改为向量化版本:

# 预定义常量,减少重复计算
const inv_R = 1 / 0.001987
const inv_T0 = 1 / 273.15
const T0_val = 273.15

function w(x, dG, dH, dcp)
    exp.(-inv_R .* (dG * inv_T0 .+ dH .* (1 ./ x .- inv_T0) .+ dcp .* (1 .- T0_val ./ x .- log.(x ./ T0_val))))
end

function dh(x, dH, dcp)
    dH .+ dcp .* (x .- T0_val)
end

function ii(v, w, dcp)
    dcp .* v^2 .* w .* (6v^2 .* w .+ 9v^2 .+ 6v .* w.^2 .+ 12v .* w .+ 12v .+ 5w.^4 .+ 8w.^3 .+ 9w.^2 .+ 8w .+ 5)
end

function hh(v, w, dh_val)
    dh_val .* v^2 .* w .* (6v^2 .* w .+ 9v^2 .+ 6v .* w.^2 .+ 12v .* w .+ 12v .+ 5w.^4 .+ 8w.^3 .+ 9w.^2 .+ 8w .+ 5)
end

function h2(v, w, dh_val)
    dh_val.^2 .* v^2 .* w .* (12v^2 .* w .+ 9v^2 .+ 18v .* w.^2 .+ 24v .* w .+ 12v .+ 25w.^4 .+ 32w.^3 .+ 27w.^2 .+ 16w .+ 5)
end

function Q(v, w)
    3v^5 .+ 3v^4 .* w.^2 .+ 9v^4 .* w .+ 19v^4 .+ 2v^3 .* w.^3 .+ 6v^3 .* w.^2 .+ 12v^3 .* w .+ 30v^3 .+ v^2 .* w.^5 .+ 2v^2 .* w.^4 .+ 3v^2 .* w.^3 .+ 4v^2 .* w.^2 .+ 5v^2 .* w .+ 21v^2 .+ 7v .+ 1
end

2. 重构模型,避免循环内重复计算

原模型逐点计算中间变量,存在大量冗余函数调用。将所有计算移到循环外,一次性完成向量化运算,并用多变量正态分布替代循环单变量似然计算:

@model function model_cp(x, y)
    σ ~ truncated(Normal(0,10), 0, Inf)
    v ~ Normal(0.048,0.06)
    dG ~ Normal(-0.2,0.2)
    dH ~ Normal(-1,0.5)
    dcp ~ Normal(0.01,0.03)

    # 一次性计算所有样本的中间变量
    w_val = w(x, dG, dH, dcp)
    dh_val = dh(x, dH, dcp)
    Q_val = Q(v, w_val)
    ii_val = ii(v, w_val, dcp)
    hh_val = hh(v, w_val, dh_val)
    h2_val = h2(v, w_val, dh_val)

    μ = (ii_val .+ (hh_val .+ h2_val.^2) ./ (inv_R .* x.^2)) ./ Q_val

    # 用MvNormal批量计算似然,替代循环内的Normal
    y ~ MvNormal(μ, σ * I)
end

3. 优化自动微分(AD)后端

NUTS的采样速度高度依赖AD效率,Turing默认使用ForwardDiff,可切换到ReverseDiff提升复杂模型的AD速度:

using ReverseDiff
Turing.setadbackend(:reversediff)
ReverseDiff.compilecache(true)  # 开启编译缓存,进一步加速

若ReverseDiff兼容性出现问题,可切换回ForwardDiff并调整chunk大小:

using ForwardDiff
Turing.setadbackend(:forwarddiff)
ForwardDiff.chunk_size = 10  # 根据参数数量调整chunk值

4. 调整NUTS采样参数

适当调整NUTS的适应步数和目标接受率,提升采样稳定性与效率:

# 增加适应步数到2000,目标接受率设为0.85(默认0.8)
chains = sample(model, NUTS(2000, 0.85), 30000; burnin=6000)

原代码(修改前)

# Load necessary libraries
using Turing, MCMCChains, Random,Plots,StatsPlots, Statistics

# Generated function used to evaluate model
@generated function w_generated(x, dG, dH, dcp)
    quote
        return exp(-1 / R * (dG / T0 + dH * (1 / x - 1 / T0) + dcp * (1 - T0 / x - log(x / T0))))
    end
end

@generated function dh_generated(x, dH, dcp)
    quote
        return dH + dcp * (x - T0)
    end
end

@generated function ii_generated(v, w, dcp)
    quote
        return dcp * v^2 * w * (6 * v^2 * w + 9 * v^2 + 6 * v * w^2 + 12 * v * w + 12 * v + 5 * w^4 + 8 * w^3 + 9 * w^2 + 8 * w + 5)
    end
end

@generated function hh_generated(v, w, dh)
    quote
        return dh * v^2 * w * (6 * v^2 * w + 9 * v^2 + 6 * v * w^2 + 12 * v * w + 12 * v + 5 * w^4 + 8 * w^3 + 9 * w^2 + 8 * w + 5)
    end
end

@generated function h2_generated(v, w, dh)
    quote
        return dh^2 * v^2 * w * (12 * v^2 * w + 9 * v^2 + 18 * v * w^2 + 24 * v * w + 12 * v + 25 * w^4 + 32 * w^3 + 27 * w^2 + 16 * w + 5)
    end
end

@generated function Q_generated(v, w)
    quote
        return 3 * v^5 + 3 * v^4 * w^2 + 9 * v^4 * w + 19 * v^4 + 2 * v^3 * w^3 + 6 * v^3 * w^2 + 12 * v^3 * w + 30 * v^3 + v^2 * w^5 + 2 * v^2 * w^4 + 3 * v^2 * w^3 + 4 * v^2 * w^2 + 5 * v^2 * w + 21 * v^2 + 7 * v + 1
    end
end

# Generate some artificial data with model functions
Random.seed!(12)
N = 101 # number of data points
R = 0.001987 # constant R
T0 = 273.15 # constant T0

x = collect(273.15:1:373.15) # temperature values in Kelvin from 0 to 100°C

# True parameter values
v_true = 0.048
dG_true = -0.24
dH_true = -1.0
dcp_true = -0.01
σ_true = 0.001 # added noise

# Generate data using these functions and the true parameter values
w_true = w_generated.(x, dG_true, dH_true, dcp_true)
dh_true = dh_generated.(x, dH_true, dcp_true)
Q_true = Q_generated.(v_true, w_true)
ii_true = ii_generated.(v_true, w_true, dcp_true)
hh_true = hh_generated.(v_true, w_true, dh_true)
h2_true = h2_generated.(v_true, w_true, dh_true)

#SIMULATED data
y = (ii_true .+ (hh_true + h2_true.^2) ./ (R .* x.^2)) ./ Q_true .+ randn(N) .* σ_true


# Define the polynomial regression model
@model function model_cp(x, y, N)
    σ ~ truncated(Normal(0,10), 0, Inf) # smaller noise
    v ~ Normal(0.048,0.06)
    dG ~ Normal(-0.2,0.2)
    dH ~ Normal(-1,0.5)
    dcp ~ Normal(0.01,0.03)
    for n in 1:N
        w = w_generated(x[n], dG, dH, dcp)
        dh = dh_generated(x[n], dH, dcp)
        Q = Q_generated(v, w)
        ii = ii_generated(v, w, dcp)
        hh = hh_generated(v, w, dh)
        h2 = h2_generated(v, w, dh)

        μ = (ii + (hh + h2^2) / (R .* x[n]^2)) / Q

        y[n] ~ Normal(μ, σ)
    end
end


# Perform MCMC inference, NUTS sampler
model = model_cp(x, y, N)
chains = sample(model, NUTS(),  30000, burn_in=6000)


# Plot the results
p = plot(chains) # MCMC diagnostics plots
savefig(p, "mcmc_diagnostics.png") # Save the diagnostics plot

# Plot the data with the best fit line
v_hat = mean(chains[:v])  # estimated v
dG_hat = mean(chains[:dG])  # estimated dG
dH_hat = mean(chains[:dH])  # estimated dH
dcp_hat = mean(chains[:dcp])  # estimated dcp

w_hat = w_generated.(x, dG_hat, dH_hat, dcp_hat)

dh_hat = dh_generated.(x, dH_hat, dcp_hat)
Q_hat = Q_generated.(v_hat, w_hat)
ii_hat = ii_generated.(v_hat, w_hat, dcp_hat)
hh_hat = hh_generated.(v_hat, w_hat, dh_hat)
h2_hat = h2_generated.(v_hat, w_hat, dh_hat)

y_hat = (ii_hat .+ (hh_hat + h2_hat.^2) ./ (R .* x.^2)) ./ Q_hat

p = scatter(x, y, label="Data")  # plot the original data
plot!(x, y_hat, label="Best Fit Line")  # add the best fit line
savefig(p, "data_and_best_fit.png")  # Save the data and best fit plot

内容的提问来源于stack exchange,提问作者RuBisCO

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 00:52:44