使用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
相关产品推荐
相关产品推荐

