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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 16:05:33