使用DiffEqFlux的sciml_train训练NeuralODE时触发MethodError报错
问题背景
运行DiffEqFlux官方神经ODE示例代码时抛出方法匹配错误,此前使用近似代码可正常运行,当前运行环境为Julia v1.7.1,所用代码取自DiffEqFlux v1.13.0官方教程。
复现代码
using DiffEqFlux, OrdinaryDiffEq, Flux, Optim, Plots u0 = Float32[2.0; 0.0] datasize = 30 tspan = (0.0f0, 1.5f0) tsteps = range(tspan[1], tspan[2], length = datasize) function trueODEfunc(du, u, p, t) true_A = [-0.1 2.0; -2.0 -0.1] du .= ((u.^3)'true_A)' end prob_trueode = ODEProblem(trueODEfunc, u0, tspan) ode_data = Array(solve(prob_trueode, Tsit5(), saveat = tsteps)) dudt2 = FastChain((x, p) -> x.^3, FastDense(2, 50, tanh), FastDense(50, 2)) prob_neuralode = NeuralODE(dudt2, tspan, Tsit5(), saveat = tsteps) function predict_neuralode(p) Array(prob_neuralode(u0, p)) end function loss_neuralode(p) pred = predict_neuralode(p) loss = sum(abs2, ode_data .- pred) return loss, pred end # 训练过程回调函数 list_plots = [] iter = 0 callback = function ( l, pred; doplot = false) global list_plots, iter if iter == 0 list_plots = [] end iter += 1 display(l) # 绘制当前预测值与真实数据对比图 plt = scatter(tsteps, ode_data[1,:], label = "data") scatter!(plt, tsteps, pred[1,:], label = "prediction") push!(list_plots, plt) if doplot display(plot(plt)) end return false end result_neuralode = DiffEqFlux.sciml_train(loss_neuralode, prob_neuralode.p, ADAM(0.05), cb = callback, maxiters = 300)
报错信息
MethodError: no method matching (OptimizationFunction{false, GalacticOptim.AutoZygote, OptimizationFunction{true, GalacticOptim.AutoZygote, DiffEqFlux.var"#84#89"{typeof(loss_neuralode)}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing}, GalacticOptim.var"#268#278"{GalacticOptim.var"#267#277"{OptimizationFunction{true, GalacticOptim.AutoZygote, DiffEqFlux.var"#84#89"{typeof(loss_neuralode)}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing}, Nothing}}, GalacticOptim.var"#271#281"{GalacticOptim.var"#267#277"{OptimizationFunction{true, GalacticOptim.AutoZygote, DiffEqFlux.var"#84#89"{typeof(loss_neuralode)}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing}, Nothing}}, GalacticOptim.var"#276#286", Nothing, Nothing, Nothing})(::OptimizationFunction{true, GalacticOptim.AutoZygote, DiffEqFlux.var"#84#89"{typeof(loss_neuralode)}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing}, ::GalacticOptim.AutoZygote, ::GalacticOptim.var"#268#278"{GalacticOptim.var"#267#277"{OptimizationFunction{true, GalacticOptim.AutoZygote, DiffEqFlux.var"#84#89"{typeof(loss_neuralode)}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing}, Nothing}}, ::GalacticOptim.var"#271#281"{GalacticOptim.var"#267#277"{OptimizationFunction{true, GalacticOptim.AutoZygote, DiffEqFlux.var"#84#89"{typeof(loss_neuralode)}, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing, Nothing}, Nothing}}, ::GalacticOptim.var"#276#286", ::Nothing, ::Nothing, ::Nothing) Stacktrace: [1] instantiate_function(f::Function, x::Vector{Float32}, ::GalacticOptim.AutoZygote, p::Nothing, num_cons::Int64) @ GalacticOptim C:\Users\User 1\.julia\packages\GalacticOptim\fow0r\src\function\zygote.jl:40 [2] instantiate_function(f::Function, x::Vector{Float32}, ::GalacticOptim.AutoZygote, p::Nothing) @ GalacticOptim C:\Users\User 1\.julia\packages\GalacticOptim\fow0r\src\function\zygote.jl:4 [3] sciml_train(::typeof(loss_neuralode), ::Vector{Float32}, ::ADAM, ::Nothing; lower_bounds::Nothing, upper_bounds::Nothing, maxiters::Int64, kwargs::Base.Pairs{Symbol, var"#43#45", Tuple{Symbol}, NamedTuple{(:cb,), Tuple{var"#43#45"}}}) @ DiffEqFlux C:\Users\User 1\.julia\packages\DiffEqFlux\gH716\src\train.jl:87
排查思路
- 先排除Julia主版本影响:Julia 1.7.1在DiffEqFlux v1.13.0的官方兼容版本范围内,主版本不兼容的可能性极低
- 从报错栈定位根因:错误触发点在
GalacticOptim包构造OptimizationFunction的环节,属于典型的依赖包版本接口不匹配问题 - 核对SciML生态迭代记录:对应版本迭代周期内SciML对优化层做过架构调整,原
GalacticOptim后续重命名为Optimization.jl,若本地安装的GalacticOptim版本与DiffEqFlux v1.13.0要求的依赖版本不一致,就会出现构造函数参数列表不匹配的MethodError - 提前排查次级适配点:部分版本迭代中调整过训练回调的入参顺序,修复版本问题后如果触发回调相关报错,可对应调整入参顺序。
解决方案
- 固定兼容的依赖版本,在Julia REPL中执行以下命令安装匹配的包版本:
using Pkg Pkg.add(name="DiffEqFlux", version="1.13.0") Pkg.add(name="GalacticOptim", version="2.0.2") Pkg.add(name="OptimizationFlux", version="0.1.0") Pkg.resolve() - 重启Julia会话后重新运行代码即可正常执行。
- 如果不想固定旧版本,也可以直接升级DiffEqFlux到最新稳定版,按照最新版规范调整
sciml_train为新的Optimization.solve接口写法,适配新版Optimization生态的调用规范。
注意:升级到最新版时,回调函数的参数顺序会变为
callback(p, l, pred),需要对应调整回调函数入参顺序,否则会出现参数不匹配的报错。
内容的提问来源于stack exchange,提问作者Adam Davis
相关产品推荐
相关产品推荐

