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

使用Julia Flux训练含输出导数自定义损失的神经网络问题咨询

问题排查与解决方案

1. 损失函数正确性判断

你当前的损失函数形式符合哈密顿系统残差约束逻辑:对应最优控制中协态方程、状态方程、最优性条件的残差平方和,形式本身没有错误。
但存在两个使用层面的问题:

  • 当前损失是单个时间点的损失,Flux.train!会逐个遍历你传入的t序列,每个时间点单独做一次梯度更新,而非对全时间区间的总损失做整体优化,容易出现训练不稳定、单点波动覆盖全局趋势的问题。
  • 你混用了ForwardDiff和Zygote两种自动微分框架:Flux默认用Zygote计算网络参数的梯度,而你用ForwardDiff计算时间维度的导数,这种嵌套微分大概率会出现梯度截断,导致网络参数收不到更新信号。

2. x值不变的原因

你遇到的x在非初始时间点保持固定的问题,来自两个可修复的bug:

  • 你的三个网络输出层都加了relu激活:如果初始权重偏小,输出层输出会恒为0,代入x(t)的定义(t - t₀)*X([t])[1] + x₀就会变成0 + x₀,自然所有时间点x都是初始值。建议先把输出层的relu去掉,或者替换为tanh等无零边界的激活函数。
  • 刚才提到的梯度截断问题:如果参数的梯度计算结果为nothing,网络参数不会有任何更新,输出自然一直不变。你可以用Zygote.gradient(() -> loss(t[1]), Θ)打印梯度验证这个问题。

3. 损失变化观测方法

你当前的回调函数(cb)只打印了固定字符串,没有输出损失值。你可以通过修改回调函数记录每次迭代的损失:

# 先定义全局数组存储损失
loss_history = []
# 改损失为全时间区间平均损失,更稳定
function total_loss()
    sum(loss(ti) for ti in t) / length(t)
end
# 训练时修改cb
cb = () -> begin
    l = total_loss()
    push!(loss_history, l)
    println("当前损失: ", l)
end
# 训练时传入总损失,而非单时间点损失
Flux.train!(total_loss, Θ, Iterators.repeated((), 1000), opt, cb=cb)

4. 修复后的核心代码调整参考

# 去掉输出层的relu
X = Chain(Dense(1,len_hidden, tanh), Dense(len_hidden,1))
Ρ = Chain(Dense(1,len_hidden, tanh), Dense(len_hidden,1))
U = Chain(Dense(1,len_hidden, tanh), Dense(len_hidden,1))

# 改用Zygote计算时间导数,避免嵌套微分冲突
dxdt(t) = Zygote.gradient(x, t)[1]
dpdt(t) = Zygote.gradient(p, t)[1]

内容的提问来源于stack exchange,提问作者Gabriel RM

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 09:27:03