对TensorFlow数值积分结果求梯度时如何避免矩阵求逆错误
我使用tfp.math.ode.BDF对常微分方程(ODE)系统进行数值积分,与API文档中的示例代码逻辑一致,我通过ode_fn(t, y, theta)函数定义待求解的ODE系统,目前可正常求取ode_fn关于theta的梯度,也可正常使用tfp.math.ode.BDF完成ODE积分运算。
但当我尝试对ODE求解结果求取关于theta的梯度时,出现如下报错。当我将ode_fn替换为更简单的ODE集合时,代码可无异常运行,请问是否需要调整求解器设置以规避该错误?
InvalidArgumentError Traceback (most recent call last) <ipython-input-9-77ebcb7dd888> in <module>() ----> 1 print(g.gradient(foo, theta0)) 5 frames /usr/local/lib/python3.7/dist-packages/tensorflow/python/eager/backprop.py in gradient(self, target, sources, output_gradients, unconnected_gradients) 1088 output_gradients=output_gradients, 1089 sources_raw=flat_sources_raw, -> 1090 unconnected_gradients=unconnected_gradients) 1091 1092 if not self._persistent: /usr/local/lib/python3.7/dist-packages/tensorflow/python/eager/imperative_grad.py in imperative_grad(tape, target, sources, output_gradients, sources_raw, unconnected_gradients) 75 output_gradients, 76 sources_raw, ---> 77 compat.as_str(unconnected_gradients.value)) /usr/local/lib/python3.7/dist-packages/tensorflow/python/eager/function.py in _backward_function_wrapper(*args) 1301 break 1302 return backward._call_flat( # pylint: disable=protected-access -> 1303 processed_args, remapped_captures) 1304 1305 return _backward_function_wrapper, recorded_outputs /usr/local/lib/python3.7/dist-packages/tensorflow/python/eager/function.py in _call_flat(self, args, captured_inputs, cancellation_manager) 1962 # No tape is watching; skip to running the function. 1963 return self._build_call_outputs(self._inference_function.call( -> 1964 ctx, args, cancellation_manager=cancellation_manager)) 1965 forward_backward = self._select_forward_and_backward_functions( 1966 args, /usr/local/lib/python3.7/dist-packages/tensorflow/python/eager/function.py in call(self, ctx, args, cancellation_manager) 594 inputs=args, 595 attrs=attrs, ---> 596 ctx=ctx) 597 else: 598 outputs = execute.execute_with_cancellation( /usr/local/lib/python3.7/dist-packages/tensorflow/python/eager/execute.py in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name) 58 ctx.ensure_initialized() 59 tensors = pywrap_tfe.TFE_Py_Execute(ctx._handle, device_name, op_name, ---> 60 inputs, attrs, num_outputs) 61 except core._NotOkStatusException as e: 62 if name is not None: InvalidArgumentError: Input matrix is not invertible. [[{{node gradients/IdentityN_grad/bdfGradients/while/body/_718/gradients/IdentityN_grad/bdfGradients/while/bdf/while/body/_2247/gradients/IdentityN_grad/bdfGradients/while/bdf/while/while/body/_3245/gradients/IdentityN_grad/bdfGradients/while/bdf/while/while/while/body/_5670/gradients/IdentityN_grad/bdfGradients/while/bdf/while/while/while/while/body/_7588/gradients/IdentityN_grad/bdfGradients/while/bdf/while/while/while/while/triangular_solve/MatrixTriangularSolve}}]] [Op:__inference___backward_debug_ode_solver_9192_32890] Function call stack: __backward_debug_ode_solver_9192
问题根因
该Input matrix is not invertible错误触发于BDF求解器的伴随状态反向梯度计算流程:隐式BDF格式求解梯度时,需要对ODE右侧函数关于状态量y的雅克比矩阵做三角分解求逆,你自定义的ode_fn在积分过程的某个时间步对应的雅克比矩阵出现了数值奇异,无法完成求逆操作。简单ODE系统可正常运行的原因是其雅克比矩阵在全积分区间都保持满秩,不会触发该错误。
解决方案
1. 调整BDF求解器配置
- 降低求解器最高阶数:将BDF初始化参数
max_order从默认值5调低到1~2,低阶隐式格式对雅克比矩阵条件数的容忍度更高,可大幅降低奇异概率 - 收紧误差容限:将
rtol(相对误差)、atol(绝对误差)参数调低1~3个数量级,更严格的步长控制会避免求解器进入雅克比奇异的数值区间 - 手动传入解析雅克比:不要依赖框架自动生成数值雅克比,自行推导
ode_fn关于状态y的雅克比解析式,作为jacobian_fn参数传入BDF求解器,可消除自动微分引入的数值误差导致的伪奇异问题
2. 更换计算方案
- 改用显式求解器:如果你的ODE系统刚性不强,可以替换为
tfp.math.ode.DormandPrince显式RK求解器,显式求解器的反向梯度计算不需要对雅克比矩阵求逆,天然不会触发该错误 - 启用双精度计算:如果必须使用BDF求解器,可以在积分时设置
enable_double_precision=True,更高的数值精度可以缓解大部分数值奇异问题
3. 正则化ODE系统
- 给
ode_fn的输出添加1e-12量级的微小随机扰动,或者对参数theta的取值范围添加约束,避免参数取值落入会导致ODE系统出现奇点的区间
内容的提问来源于stack exchange,提问作者Jack Elsey
相关产品推荐
相关产品推荐

