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

对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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 00:27:02