基于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
相关产品推荐
相关产品推荐

