在Julia中用SciML复现螺旋示例时遇Zygote数组突变不支持错误
问题分析
错误提示Mutating arrays is not supported核心原因是Zygote无法对StatsBase.sample的原地修改操作进行微分——你的test函数在求导路径中调用了get_batch,导致Zygote试图处理采样过程中的数组突变。此外还有两个潜在问题:
- Lux模型调用时未正确处理状态(
st),可能引发后续错误 - ODE问题的初始
tspan设置为[0.,0.]无效,不符合求解要求
解决方案
1. 将采样移出求导路径
训练时采样属于非微分操作,需提前获取batch数据,再传入损失函数,避免Zygote尝试对采样过程求导:
# 修改损失函数,接收外部传入的batch数据 function test(θ, y0s, ts, targets) pred = predict(θ, y0s, ts) loss = sum(abs2, targets .- pred) end # 提前获取batch,将其移出求导路径 y0s, ts, targets = get_batch() # 对固定batch的损失函数求导 x, lambda = pullback((θ) -> test(θ, y0s, ts, targets), p_init) lambda(x) # 不再触发数组突变错误
2. 修复Lux模型的状态处理
Lux模型调用会返回更新后的状态,需用StatefulLux自动管理状态,避免手动传递时的错误:
using Lux.Experimental: StatefulLux # 用StatefulLux封装模型与初始状态 neural_net_stateful = StatefulLux.StatefulLux(neural_net, st) function neural_net_func!(du, u, p, t) du .= neural_net_stateful(u.^3, p)[1] end
3. 修正ODE问题的初始tspan
将无效的tspan = [0.,0.]改为合理的时间区间,与真实问题保持一致:
prob_nn = ODEProblem(neural_net_func!, [0., 0.], (0.0, 25.0), p)
完整修正代码
using Lux, DiffEqFlux, DifferentialEquations, ComponentArrays, Random, StatsBase, MLUtils using Zygote using Lux.Experimental: StatefulLux true_y0 = [2., 0.] true_A = [-0.1 2.; -2. -0.1] data_size = 1000 times = LinRange(0, 25, data_size) function ground_truth!(du, u, p, t) du .= true_A * (u.^3) end ground_truth_odeProb = ODEProblem(ground_truth!, true_y0, (0, times[end])) sol_ode = Array(solve(ground_truth_odeProb, Tsit5(), abstol=1e-10, reltol=1e-10, saveat=times)) batch_time = 10 batch_size = 20 function get_batch() s = sample(range(1, data_size - batch_time), batch_size, replace=false) batch_y0 = sol_ode[:, s] batch_t = times[1:batch_time] batch_y = stack([sol_ode[:, s .+ i] for i in 1:batch_time], dims=3) return batch_y0, batch_t, batch_y end const neural_net = Lux.Chain(Lux.Dense(2, 50, tanh), Lux.Dense(50, 2)) rng = Random.default_rng() p, st = Lux.setup(rng, neural_net) p_init = ComponentArray(p) # 用StatefulLux管理模型状态 neural_net_stateful = StatefulLux.StatefulLux(neural_net, st) function neural_net_func!(du, u, p, t) du .= neural_net_stateful(u.^3, p)[1] end # 修正初始时间区间 prob_nn = ODEProblem(neural_net_func!, [0., 0.], (0.0, 25.0), p) function predict(θ, y0s, ts) _prob = remake(prob_nn, u0=y0s, tspan=(ts[1], ts[end]), p=θ) Array(solve(_prob, Tsit5(), saveat=ts, abstol=1e-5, reltol=1e-5)) end # 损失函数接收外部传入的batch数据 function test(θ, y0s, ts, targets) pred = predict(θ, y0s, ts) sum(abs2, targets .- pred) end # 提前获取batch,移出求导路径 y0s, ts, targets = get_batch() x, lambda = pullback((θ) -> test(θ, y0s, ts, targets), p_init) lambda(x) # 正常运行
关键说明
- Zygote对原地修改操作严格限制,所有涉及数组突变的操作(如部分采样、原地赋值)必须移出求导路径。
StatefulLux是Lux官方提供的状态管理工具,可自动跟踪模型状态变化,避免手动传递状态的错误。- ODE问题的
tspan必须是有效时间区间,否则求解器会直接报错。
内容的提问来源于stack exchange,提问作者user22881792
相关产品推荐
相关产品推荐

