Julia/Turing采样带强迫项微分方程时出现TypeError报错
带强迫项ODE拟合Turing采样类型错误解决
问题表现
拟合带插值强迫项的0维箱式常微分方程时,单独调用DifferentialEquations求解器可正常运行;接入Turing框架使用NUTS/HMC采样做贝叶斯参数估计时,梯度计算环节触发如下类型错误:TypeError: in typeassert, expected Float64, got a value of type ForwardDiff.Dual{Nothing, Float64, 3}
报错复现代码
using Interpolations, DifferentialEquations, Plots, Turing, LinearAlgebra # 定义ODE右端函数 function f(du,u,p,t) α, β, F = p du[1] = α * u[1] + β * F(t) # Interpolations.jl通过()调用完成插值 end # 定义初值问题 begin # 定义强迫项 time_forcing = -1.:9. data_forcing = [10,0,0,0,0,0,0,0,0, 0, 0] F = Interpolations.scale(interpolate(data_forcing, BSpline(Linear())), time_forcing) α = -0.5 β = 1 p_lin = (α, β, F) # 定义初值与时间跨度 u0 = [0.] tspan = (-1.,9.) ode_lin = ODEProblem(f,u0,tspan,p_lin) end # 生成带噪声模拟数据 sol = solve(ode_lin, Tsit5(); saveat=1) data = Array(sol) + 0.2 * randn(size(Array(sol))) @model function fit_simple_box(data, F, ode_lin) # 定义先验分布 σ ~ InverseGamma(2, 3) α ~ Normal(0, 3) β ~ Normal(0, 3) # 求解ODE p = (α, β, F) predicted = solve(ode_lin, Tsit5(); p=p, saveat=1) # 定义似然 for i in 1:length(predicted) data[i] ~ Normal(predicted[i][1], σ^2) end return nothing end model = fit_simple_box(data, F, ode_lin) chain = sample(model, NUTS(0.65), MCMCSerial(), 1000, 2)
问题根因
报错来自三个类型不兼容问题:
- 模型外预构造的
ode_lin是固定绑定Float64数值类型的ODE问题实例,Turing调用ForwardDiff执行自动微分时,会传入ForwardDiff.Dual类型的待估参数α、β,直接复用固定类型的问题实例无法完成自动类型提升。 - 插值对象F是固定不变的非数值参数,被放入待微分的参数元组p中时,自动微分引擎会错误地对F执行类型追踪和微分操作,直接触发类型断言失败。
- 原生构造的插值对象默认返回
Float64类型结果,与Dual类型的待估参数做运算时无适配逻辑,进一步触发类型不匹配错误。
修复方案
通过三点调整即可解决报错:
- 不要在模型外部预构造固定数值类型的ODEProblem,将ODE右端函数、初值、时间跨度定义为全局通用对象,在模型内部结合当前采样得到的参数构造对应类型的问题实例,保证类型可自动适配- 将固定不变的强迫项插值对象声明为全局常量,不要放入待微分的参数序列中,避免被自动微分引擎错误追踪。
- 调用solve时显式指定前向自动微分灵敏度算法,开启时间跨度类型转换配置,避免自动微分逻辑误判参数属性。
修复后可运行代码
using Interpolations, DifferentialEquations, Plots, Turing, LinearAlgebra, SciMLSensitivity # 定义ODE右端 function f(du,u,p,t) α, β = p du[1] = α * u[1] + β * F_const(t) end # 定义固定强迫项,声明为常量避免被自动微分引擎追踪 const time_forcing = -1.:9. const data_forcing = [10,0,0,0,0,0,0,0,0, 0, 0] const F_const = Interpolations.scale( interpolate(data_forcing, BSpline(Linear())), time_forcing ) # 生成模拟数据 begin α_true = -0.5 β_true = 1 p_true = (α_true, β_true) u0 = [0.] tspan = (-1.,9.) ode_prob = ODEProblem(f,u0,tspan,p_true) sol = solve(ode_prob, Tsit5(); saveat=1) data = Array(sol) + 0.2 * randn(size(Array(sol))) end @model function fit_simple_box(data) # 先验分布 σ ~ InverseGamma(2, 3) α ~ Normal(0, 3) β ~ Normal(0, 3) # 构造当前参数对应的问题实例,保证类型适配 p = (α, β) u0_cur = [zero(α)] prob = ODEProblem(f, u0_cur, (-1.0, 9.0), p) predicted = solve( prob, Tsit5(); saveat=1, sensealg=ForwardDiffSensitivity(convert_tspan=true) ) # 似然 for i in 1:length(predicted) data[i] ~ Normal(predicted[i][1], σ^2) end return nothing end model = fit_simple_box(data) chain = sample(model, NUTS(0.65), MCMCSerial(), 1000, 2)
若使用较旧版本的DifferentialEquations生态,可将
SciMLSensitivity替换为DiffEqSensitivity即可正常运行。
内容的提问来源于stack exchange,提问作者Duo Chan
相关产品推荐
相关产品推荐

