GradRepair-ODE:神经网络ODE训练中的认证梯度修复
摘要
GradRepair-ODE引入了一个可靠性框架,用于认证和修复神经网络ODE训练中的梯度,以解决科学机器学习和生成模型中的数值稳定性问题。
arXiv:2609.13204v1 公告类型:新
摘要:神经网络ODE在训练循环中使用数值求解器。求解器决定了前向轨迹,并影响传递给优化器的梯度。这种耦合给科学机器学习和连续时间生成建模带来了可靠性问题,包括扩散概率流常微分方程和流匹配模型。在宽松步长、刚性动力学、混沌敏感性或事件不连续性下,可微分ODE管道可能返回有限梯度,但其方向在数值上可疑。我们引入GradRepair-ODE,这是一个在优化器步骤中检查、修复和拒绝ODE梯度的可靠性框架。该方法计算多个梯度候选,通过方向有限差分检查和求解器诊断进行比较,诊断可能的数值故障模式,通过路径切换或更严格的重新计算修复选定的梯度,并拒绝下降方向无法认证的步骤。在六个合成ODE系统中,GradRepair-ODE保持低风险系统不变,将Robertson和Lorenz梯度修复到与严格参考的余弦相似度为1.000,将不安全接受步骤从37减少到0,并拒绝了一个事件不连续的情况,而不是应用未经认证的更新。本文主张对训练契约进行简单更改:ODE梯度应附带数值证据到达优化器。
查看缓存全文
缓存时间: 2026/09/15 08:36
# GradRepair-ODE:用于神经ODE训练的经过认证的梯度修复
来源:https://arxiv.org/html/2609.13204 \[ BoldFont = texgyretermes\-bold\.otf, ItalicFont = texgyretermes\-italic\.otf, BoldItalicFont = texgyretermes\-bolditalic\.otf \]
###### 摘要
神经常微分方程在训练循环内使用数值求解器。求解器决定了前向轨迹,也影响了传递给优化器的梯度。这种耦合为科学机器学习和连续时间生成建模(包括扩散概率流常微分方程和流匹配模型)带来了可靠性问题。在宽松的步长、刚性动力学、混沌敏感性或事件不连续性下,可微分的ODE流程可能返回有限梯度,但其方向在数值上是可疑的。我们提出了GradRepair-ODE,这是一个在优化器步骤中检查、修复和拒绝ODE梯度的可靠性框架。该方法计算多个梯度候选,通过方向有限差分检查和求解器诊断进行比较,诊断可能的数值失效模式,通过路径切换或更严格的重计算修复选定的梯度,并拒绝无法证明下降方向的步骤。在六个合成ODE系统中,GradRepair-ODE对低风险系统保持不变,将Robertson和Lorenz梯度的余弦相似度修复到1.000±0.001(相对于严格参考),将不安全接受步骤从37减少到0,并拒绝了一个事件不连续的情况,而不是应用未经验证的更新。本文主张对训练合约进行一个简单的改变:ODE梯度到达优化器时应附带数值证据。
简短标题:GradRepair-ODE
作者:Z. Bi 和 X. L. Chia
关键词:.神经常微分方程;可微分模拟;伴随敏感性;数值稳定性;梯度可靠性;科学机器学习;扩散模型;流匹配;连续标准化流。
数学学科分类:.65L05;65L06;65L20;65L50;68T07;90C30。
SISC部分:.科学计算的机器学习方法。
## 1 引言
神经常微分方程(Neural ODEs)用由数值常微分方程(ODE)求解器计算轨迹的连续时间动力系统取代了离散的变换堆栈 \[4 (https://arxiv.org/html/2609.13204#bib.bib1)\]。这种表述赋予了模型一个吸引人的数学接口,并允许求解器容差以数值精度换取速度。它还将训练的大部分内容移入数值分析:优化器接收到的不是一个精确梯度,而是一个被求解器误差、插值、刚性、时间范围长度和所选敏感性路径所塑造的梯度。现代可微分ODE方法既暴露了通过求解器操作的直接反向传播,也暴露了具有不同内存和精度权衡的伴随方法 \[4 (https://arxiv.org/html/2609.13204#bib.bib1),14 (https://arxiv.org/html/2609.13204#bib.bib14)\]。相同的数值问题超出了经典的神经ODE基准测试范围。连续标准化流(CNFs)使用ODE动态进行密度建模 \[7 (https://arxiv.org/html/2609.13204#bib.bib2)\];扩散概率模型和基于分数的随机微分方程(SDE)模型将生成建模与连续时间动态和概率流ODE联系起来 \[23 (https://arxiv.org/html/2609.13204#bib.bib3),13 (https://arxiv.org/html/2609.13204#bib.bib4)\];流匹配通过沿概率路径回归向量场来训练CNFs,并明确地将扩散路径作为一个重要情况包含在内 \[15 (https://arxiv.org/html/2609.13204#bib.bib5)\];整流流学习ODE传输,旨在使生成路径变直并提高求解器效率 \[16 (https://arxiv.org/html/2609.13204#bib.bib6)\]。在所有这些设置中,学习到的向量场最终都与数值积分配对。当通过可微分求解训练此类模型或通过基于ODE的似然或采样例程评估它们时,梯度可靠性不再是一个小众的实现细节;它是连续时间生成建模的数值基础的一部分。先前的工作表明,先优化后离散化的伴随和先离散后优化的梯度可能以对神经ODE训练重要的方式不一致 \[6 (https://arxiv.org/html/2609.13204#bib.bib13),18 (https://arxiv.org/html/2609.13204#bib.bib15)\]。经典的数值分析长期以来将步长控制、密集输出、刚性、向后微分格式、有限精度效应和敏感性分析视为求解器层面的问题 \[5 (https://arxiv.org/html/2609.13204#bib.bib7),9 (https://arxiv.org/html/2609.13204#bib.bib8),22 (https://arxiv.org/html/2609.13204#bib.bib9),2 (https://arxiv.org/html/2609.13204#bib.bib10),11 (https://arxiv.org/html/2609.13204#bib.bib11),12 (https://arxiv.org/html/2609.13204#bib.bib17),10 (https://arxiv.org/html/2609.13204#bib.bib19),3 (https://arxiv.org/html/2609.13204#bib.bib20)\]。通用科学计算库同样区分非刚性显式求解器和感知刚性的隐式方法,并警告说发散或异常多的迭代可能表明刚性 \[20 (https://arxiv.org/html/2609.13204#bib.bib12)\]。这些事实在求解器层面很容易理解,但优化器接口通常只接收结果张量。一旦自动微分返回具有预期形状的有限数组,大多数训练循环几乎没有关于其方向可靠性的额外证据。这种失败很容易被忽视。一个梯度张量可以是有限的、形状正确的,但仍然指向错误的方向。如果该方向与其他敏感性路径或局部方向检查不一致,优化器可能在增加损失,而训练循环只记录正常的更新。相关问题是当前的数值证据是否支持下一步。GradRepair-ODE在优化器边界回答了这个问题。它计算多个梯度候选,从路径不一致性和有限差分残差构建可靠性证书,诊断数值风险,修复选定的梯度,并证明最终步骤是否被接受。该方法不提出另一种ODE架构或通用求解器。它添加了一个缺失的接口:梯度信任、修复和拒绝成为被记录的训练事件。其意义比调试工具更深刻。求解器容差、伴随方法和损失曲线并不能说明下一步更新是否合理。GradRepair-ODE在数值近似变为优化动作的位置放置了一个证书。该证书可以信任一个良性的梯度,将一个可修复的梯度路由到更安全的路径,或者拒绝一个平滑梯度证据已崩溃的更新。这在科学计算中很重要,因为具有破坏性的失败往往是无声的:计算以一个看似合理的数字继续进行,而该数字的方向已失去数值意义。实验使用受控机制而非装饰性应用。该套件包含平滑、刚性、混沌、事件驱动和神经向量场动力学,因此相同的可靠性层面临着它声称要暴露的病态情况。主要结果是将隐藏的优化器风险转化为明确的决策。GradRepair-ODE对低风险系统保持不变,将Robertson和Lorenz梯度的余弦相似度修复到1.000±0.001(相对于严格参考),将不安全接受步骤从37减少到0,并拒绝了事件不连续的情况,而不是应用未经证实的更新。一个有用的可靠性方法应该恰好具有这种特性:保持良性步骤廉价,当方向可以恢复时为准确性付费,并在微分对象不再与优化轨迹匹配时拒绝更新。本文贡献了四个方面。它为可微分ODE训练定义了一个面向优化器的梯度可靠性证书,使用敏感性路径不一致、方向有限差分和求解器诊断。它将修复表述为梯度路径之间的约束选择,而不是全面转向最严格的计算。它推导了一个局部下降证书,将经验梯度误差半径与下一步优化器步骤联系起来。它还报告了一项受控数值研究,包括梯度裁剪和Armijo型线搜索作为安全措施,展示了证书何时信任、修复或拒绝。结果不是训练稳定性的轶事。它是更新流的逐步数值记录。
## 2 方法
### 2.1 问题设置
我们考虑一个由参数θ参数化的ODE,\\frac\{dx\}\{dt\}=f\(t,x,\\theta\),\\qquad x\(t\_\{0\}\)=x\_\{0\},\(1\)
具有轨迹级损失 L\(\\theta\)=\\ell\(x\(t\_\{0\}:t\_\{1\};\\theta\),y\)\.\(2\)
训练循环接收到一个由数值求解器和敏感性路径产生的近似梯度 g^\\hat\{g\}。GradRepair-ODE测试 g^\\hat\{g\} 是否能够支持更新 −ηg^\\-\\eta\\hat\{g\}。导数和更新之间的区别是实际的,而非语义的。一个敏感性路径可以产生参数空间向量,即使用于构建它的轨迹是脆弱的。令 S\\mathcal\{S\} 表示求解器配置,包括积分规则、容差、步长策略、插值规则和事件处理。令 P\\mathcal\{P\} 表示敏感性路径。返回的梯度最好写成 g^=G\(θ,S,P\),\\hat\{g\}=G\(\\theta;\\mathcal\{S\},\\mathcal\{P\}\)\,\(3\)
而不是简单地 ∇L\(θ\)\\nabla L\(\\theta\)。GradRepair-ODE将 G\(θ,S,P\)G\(\\theta;\\mathcal\{S\},\\mathcal\{P\}\) 与 ∇L\(θ\)\\nabla L\(\\theta\) 之间的差距视为一个可观察的数值量。该框架以与线搜索是优化器无关的相同方式求解器无关:它保留底层数值方法,并改变优化器信任其输出的条件。每个训练步骤产生梯度路径证据、方向证据和求解器证据。梯度路径证据比较同一导数的独立或部分独立的近似。方向证据将内积与局部中心有限差分进行比较。求解器证据记录数值压力,如步长崩溃、高函数评估计数、失败步骤或已知不连续性。由这些信号支持的梯度通过成本低。未通过这些信号的梯度将被修复或扣留。
### 2.2 梯度候选
每个诊断步骤将默认的粗略求解器路径梯度与更精确的重计算路径的梯度进行比较。一个更精细的离散路径测试方向是否在减少的积分误差下存活。一个检查点式重计算路径测试更大的轨迹一致性是否改变梯度。一个严格路径为认证或最后手段修复提供高成本参考。该设计借鉴了反向模式微分、伴随敏感性分析和检查点技术中熟悉的内存-计算权衡 \[8 (https://arxiv.org/html/2609.13204#bib.bib16),12 (https://arxiv.org/html/2609.13204#bib.bib17),21 (https://arxiv.org/html/2609.13204#bib.bib18)\]。这里的目的是更新决策:在参数更改之前接受、修复或拒绝。候选集按成本和预期的数值保真度排序: C=\{gcoarse,gdisc,gckpt,gstrict\}\.\\mathcal\{C\}=\\\{g\_\{\\mathrm\{coarse\}\},g\_\{\\mathrm\{disc\}\},g\_\{\\mathrm\{ckpt\}\},g\_\{\\mathrm\{strict\}\}\\\}\.\(4\)
粗略路径是标准训练循环会消耗的梯度。离散路径以更小的步长重新计算求解,并检查方向是否在精炼中存活。检查点式路径代表更轨迹一致的重计算,具有额外成本。严格路径在受控研究中用作修复目标或评估参考。在更大的实现中,这些符号可能对应于连续伴随、通过求解器操作的直接微分、检查点离散伴随、直接敏感性方程或感知刚性的求解器。该声明不依赖于特定求解器;它依赖于多条数值路径,其不一致性携带信息。
### 2.3 可靠性证书
对于候选梯度 g_i g_\{i\},GradRepair-ODE计算成对余弦不一致 Dcos\(gi,gj\)=1−⟨gi,gj⟩‖gi‖‖gj‖\+δ,D\_\{\\cos\}\(g\_\{i\},g\_\{j\}\)=1\-\\frac\{\\langle g\_\{i\},g\_\{j\}\\rangle\}\{\\\|g\_\{i\}\\\|\\\|g\_\{j\}\\\|\+\\delta\},\(5\)
相对范数不一致 Dnorm\(gi,gj\)=‖gi−gj‖‖gi‖\+‖gj‖\+δ,D\_\{\\mathrm\{norm\}\}\(g\_\{i\},g\_\{j\}\)=\\frac\{\\\|g\_\{i\}\-g\_\{j\}\\\|\}\{\\\|g\_\{i\}\\\|\+\\\|g\_\{j\}\\\|\+\\delta\},\(6\)
以及方向有限差分(FD)残差 RFD\(g,v\)=|⟨g,v⟩−FD\(θ,v\)|\|FD\(θ,v\)|\+δ\.R\_\{\\mathrm\{FD\}\}\(g,v\)=\\frac\{\|\\langle g,v\\rangle\-\\mathrm\{FD\}\(\\theta,v\)\|\}\{\|\\mathrm\{FD\}\(\\theta,v\)\|\+\\delta\}\.\(7\)
这些项与求解器不稳定性和刚性代理结合成经验不确定性分数 ε^i=wcosDcos\+wnormDnorm\+wFDRFD\+wsolverS\+wstiffK,\\widehat\{\\epsilon\}\_\{i\}=w\_\{\\cos\}D\_\{\\cos\}\+w\_\{\\mathrm\{norm\}\}D\_\{\\mathrm\{norm\}\}\+w\_\{\\mathrm\{FD\}\}R\_\{\\mathrm\{FD\}\}\+w\_\{\\mathrm\{solver\}\}S\+w\_\{\\mathrm\{stiff\}\}K,\(8\)
其中 SS 表示求解器不稳定证据,KK 表示刚性证据。该证书分配四种状态之一:可信、可修复、不安全或失败。该证书是一个可见的可靠性估计,而非声称的数学上界。可信梯度无需修复即可应用。可修复梯度证明了额外计算的合理性,但并非原始更新。不安全梯度在可用的修复后缺乏经过验证的下降方向。失败梯度包含非有限值、无效形状或损坏的微分状态。这些状态将警告、修复、接受和失败分开。该证书还诊断失效模式。大的有限差分残差表明方向不一致。大的成对不一致表明敏感性路径不稳定。大的求解器压力表明在检查梯度之前就存在数值困难。事件标志表明平滑伴随模型可能与轨迹不匹配。当证据混合时,诊断报告数值不稳定性,而不是分配精确的物理原因。
### 2.4 修复和步骤认证
如果朴素梯度被信任,优化器可以直接使用它。如果证书报告可修复风险,GradRepair-ODE切换到更可靠的梯度路径。如果修复后的证书仍然不安全,则拒绝该步骤。在步骤层面,决策将预测的梯度信号与估计的数值不确定性进行比较: η‖g^‖2\>ηε‖g^‖\+m,\\eta\\\|\\hat\{g\}\\\|^\{2\}\>\\eta\\epsilon\\\|\\hat\{g\}\\\|\+m,\(9\)
其中 ε\\epsilon 是经验误差代理,mm 是一个小边际。该规则对于此处使用的步骤证书来说已经足够;它不是一个全局收敛定理。修复通过成本感知策略选择。当证据指向松散积分时使用容差收紧。当敏感性路径不一致但存在精炼路径时使用路径切换。对于在附加数值工作下仍然可修复的刚性或混沌情况使用严格重计算。当轨迹包含没有事件感知导数的事件不连续性,或者没有候选产生正的下降边际时,选择拒绝。该策略在可靠性约束相似文章
正交梯度约束塑造噪声标签记忆动力学
本文介绍了 OrthoGrad,这是一种几何干预,在优化过程中去除权重梯度的径向分量,并表明它在小数据场景下减少了对噪声标签的记忆,但无法防止最终的记忆。
Energy Manifold Natural Gradient Descent:面向神经PDE求解器的黎曼优化
提出了能量流形自然梯度下降(EMNGD),一种面向神经PDE求解器的流形优化框架,该框架在参数更新时与函数空间能量曲率对齐,同时遵循参数约束。理论保证和实验结果表明,该方法提高了准确性和收敛速度。
训练GNNs中的局部证据与几何读出修复
本文介绍了一种通过精确质量线性规划和事后修复分离错误原因,从而修复训练过的Graph Neural Networks中的读出的方法,并通过重新加权和集合条件翻译实现准确性提升。
流经状态:用于强化学习的神经ODE正则化
本文提出了一种基于神经ODE的正则化方法,该方法强制强化学习智能体中的潜在嵌入遵循一致的ODE流,使表示学习与环境动态对齐,并在Atari和网格世界基准上取得了性能提升。
从零开始的自动微分:PyTorch如何在物理信息神经网络中计算梯度
本文逐步追踪了PyTorch的自动微分引擎如何为物理信息神经网络训练计算梯度,包括物理残差和参数梯度所需的两个层次的微分,并使用一个简单的MLP和ODE示例。