JAX odeint求解含广义函数ODE时高阶梯度返回NaN的解决方法
如何使用JAX
odeint 生成涉及广义函数的常微分方程(ODE)的可微数值解? 问题描述
目标是得到下述初值问题的C2可微数值解:
使用JAX的初始实现代码如下:
import jax as jx import jax.numpy as jnp from jax.experimental.ode import odeint # 启用双精度 jx.config.update("jax_enable_x64", True) def calc_dy(y,x): dy = 1.0-y*y return 1.0/jnp.abs(dy) def calc_y1(y0): x_grid = jnp.linspace(0.0,1.0,2**1) y_grid = odeint(calc_dy,y0,x_grid) return y_grid[-1] calc_dy1 = jx.grad(calc_y1) calc_res = lambda y0 : 1.0/calc_dy1(y0) calc_dres = jx.grad(calc_res) y0 = jnp.array(0.0) print(f'res = {calc_res(y0)}') # res = -1.3472950405756983 print(f"dres = {calc_dres(y0)}") # dres = nan
上述代码定义的残差函数res可用于打靶法,求解使得斜率奇点恰好位于x=1.0处的临界初始条件。但当初始值高于临界值时,递归调用jx.grad得到的dres会返回nan,导致无法使用基于梯度的求解器,需要找到可行方法,让dres对任意初始值y0都返回有效计算结果。
回答
NaN产生的核心原因
二阶导返回NaN本质是两个问题叠加导致的:
- ODE右端项
1.0/jnp.abs(dy)在dy=1-y²=0处存在奇点,不仅原方程斜率趋向无穷大,abs函数在0点本身存在一阶导的跳变,其二阶导是0点处冲激形式的广义函数,直接在数值计算中调用自动微分求二阶导本身就会产生未定义值。 jax.experimental.ode.odeint默认的自适应积分器没有做奇点特殊处理,积分步长踩到奇点邻域时会产生无穷大的中间值,反向传播求高阶导时无穷大之间的运算直接生成NaN。
可落地的解决方法
1. 光滑正则化(改动成本最低)
将奇点处的非光滑项替换为处处C2可微的近似形式,从源头避免除零和导数跳变:
def calc_dy(y, x, eps=1e-7): dy = 1.0 - y*y # 用光滑的sqrt(dy²+eps)替换abs(dy),全程无尖角、无除零 return 1.0 / jnp.sqrt(dy**2 + eps)
- 正则参数
eps取1e-6~1e-8量级时,对原方程解的精度影响小于积分器本身的截断误差,同时能保证一、二阶导全程有界,不会出现NaN。 - 如果需要严格收敛到原问题的解,可以将
eps设为和积分步长同阶的小量,随网格加密同步缩小即可。 - 注意不要用
jnp.where直接屏蔽奇点邻域的分支,jnp.where在自动微分时仍会计算被屏蔽分支的数值,依然会触发除零NaN。
2. 变量替换消去奇点(精度最高)
针对这个特定形式的ODE,可以通过变量替换彻底消去右端项的奇点,不需要引入正则误差:
原方程满足$dx = |1-y^2| dy$,直接将积分变量从x替换为y,积分关系变为$x(y) = \int_{y0}^y |1-t^2| dt$,该式右端是分段光滑的多项式,任意阶导数都有明确的定义,结合隐函数定理可以直接计算res和dres,全程不需要让积分器穿过奇点,从根源上避免数值不稳定。
3. 替换支持事件处理的可微积分器(通用性最强)
如果不想修改原方程形式,可以替换为支持事件检测的可微ODE求解器:
- 提前定义奇点触发事件:当
1-y²=0时暂停积分 - 在奇点位置做局部解析延拓后再继续积分,全程保证数值解的C2连续性,避免积分步长直接踩到除零点
这种方案对任意带奇点的ODE都适用,不需要手动推导变量替换形式。
注意计算高阶梯度时需要保持双精度开启,单精度下小正则项的梯度很容易下溢触发NaN。
内容的提问来源于stack exchange,提问作者DavidJ
相关产品推荐
相关产品推荐

