Dymos/OpenMDAO最小圈速问题:纵向加速度与法向载荷耦合异常
问题描述
使用Dymos/OpenMDAO开展最小圈速问题研究时,遇到纵向加速度(du_dt)与法向载荷计算的耦合问题。尝试通过配置Newton非线性求解器和Direct线性求解器的Group耦合两个Explicit Component,但出现如下导数警告:
DerivativesWarning: Component 'traj.phases.phase0.rhs_col.cycle_group.normal_loads' has zero derivatives for the following variable pairs that were declared as 'dependent': [('fz_fl', 'du_dt'), ('fz_fr', 'du_dt'), ('fz_rl', 'du_dt'), ('fz_rr', 'du_dt')].
采用JAX进行导数计算,所有法向载荷计算均符合JAX安全规范且逻辑相对简单。请问该耦合方法是否被支持?还是必须实现Implicit Component来解决此问题?
相关代码
CombinedODE 实现
import openmdao.api as om from libs.cp_mdao_coupled_lat_lon_jax.path_constraint import TrackFollowingODE from libs.cp_mdao_coupled_lat_lon_jax.lat_lon_ode import LatLonODE from libs.cp_mdao_coupled_lat_lon_jax.curvature import Curvature from libs.cp_mdao_coupled_lat_lon_jax.normal_ode import NormalODE from cp_lapsim.track import Track from cp_config import VehicleConfig class CombinedODE(om.Group): def initialize(self): self.options.declare("num_nodes", types=int, desc="Number of nodes to be evaluated in the RHS") self.options.declare('vehicle', types=VehicleConfig, recordable=False) self.options.declare('track', types=Track, recordable=False) def setup(self): nn = self.options["num_nodes"] vehicle = self.options["vehicle"] track = self.options["track"] self.add_subsystem(name="curv", subsys=Curvature(num_nodes=nn, track=track), promotes_inputs=['s', 'u'], promotes_outputs=['ds_dt', 'kappa'] ) self.add_subsystem(name='trackfollowing', subsys=TrackFollowingODE(num_nodes=nn), promotes_inputs=['u', 'r', 'kappa'], promotes_outputs=['trackfollowing_residual'] ) cyc_group: om.Group = self.add_subsystem("cycle_group", om.Group(), promotes_inputs=['u', 'beta', 'r', 'omega', 'delta', 'motor_torque_request'], promotes_outputs=['du_dt', 'domega_dt', 'dbeta_dt', 'dr_dt', 'fz_fr', 'fz_fl', 'fz_rr', 'fz_rl', 'rear_slip_ratio'] ) cyc_group.add_subsystem(name="lat_lon", subsys=LatLonODE(num_nodes=nn, vehicle=vehicle), promotes_inputs=["u", "beta", "r", "omega", "delta", "motor_torque_request", 'fz_fr', 'fz_fl', 'fz_rr', 'fz_rl'], promotes_outputs=["du_dt", "domega_dt", "rear_slip_ratio", "dbeta_dt", "dr_dt"] ) cyc_group.add_subsystem(name="normal_loads", subsys=NormalODE(num_nodes=nn, vehicle=vehicle), promotes_inputs=["u", "du_dt", "r"], promotes_outputs=['fz_fr', 'fz_fl', 'fz_rr', 'fz_rl'] ) # Attach a Newton solver just to the subgroup cyc_group.nonlinear_solver = om.NewtonSolver(solve_subsystems=True) cyc_group.nonlinear_solver.options["maxiter"] = 10 cyc_group.nonlinear_solver.options["iprint"] = 5 cyc_group.nonlinear_solver.options["atol"] = 1e-3 cyc_group.nonlinear_solver.options["rtol"] = 1e-3 cyc_group.nonlinear_solver.linesearch = om.ArmijoGoldsteinLS() cyc_group.linear_solver = om.DirectSolver(assemble_jac=True)
NormalODE 实现
import jax import jax.numpy as jnp import openmdao.api as om from dataclasses import asdict import openmdao.jax as omj from cp_config import VehicleConfig from cp_dynamics.load_transfer import calc_steady_state_normal_loads_jax from cp_config.utils import _make_hashable class NormalODE(om.JaxExplicitComponent): def initialize(self): self.options.declare('num_nodes', types=int) self.options.declare('vehicle', types=VehicleConfig, recordable=False) def setup(self): nn = self.options['num_nodes'] # Inputs: states self.add_input('du_dt', val=jnp.full(nn, 10.0), desc='longitudinal accel', units="m/s**2") self.add_input('r', val=jnp.full(nn, 0.1), desc='yaw_rate', units="rad/s") self.add_input('u', val=jnp.full(nn, 10.0), desc='x velocity', units="m/s") # Outputs: normal loads self.add_output('fz_fr', val=jnp.full(nn, 800.0), desc='fr tire normal load', units="N") self.add_output('fz_fl', val=jnp.full(nn, 800.0), desc='fl tire normal load', units="N") self.add_output('fz_rr', val=jnp.full(nn, 800.0), desc='rr tire normal load', units="N") self.add_output('fz_rl', val=jnp.full(nn, 800.0), desc='rl tire normal load', units="N") def get_self_statics(self): vehicle = self.options['vehicle'] # Convert dataclass to dict vehicle_dict = asdict(vehicle) return _make_hashable(vehicle_dict) def compute_primal(self, du_dt, r, u): vehicle = self.options['vehicle'] nn = self.options['num_nodes'] a_x = du_dt a_y = r * u #clip extreme values u = omj.smooth_max(u, jnp.full(nn,1), mu=0.01) a_y = omj.smooth_max(a_y, jnp.full(nn,-60), mu=0.01) a_y = omj.smooth_min(a_y, jnp.full(nn, 60), mu=0.01) a_x = omj.smooth_max(a_x, jnp.full(nn, -40), mu=0.01) a_x = omj.smooth_min(a_x, jnp.full(nn, 40), mu=0.01) normal_loads, _ = calc_steady_state_normal_loads_jax(vehicle, u, a_x, a_y) # return outputs in the declared order return ( normal_loads[:, 0], # fz_fr normal_loads[:, 1], # fz_fl normal_loads[:, 2], # fz_rr normal_loads[:, 3], # fz_rl )
解答
这种用Group嵌套Newton求解器耦合Explicit Component的方法是OpenMDAO完全支持的,不需要必须改为Implicit Component。
警告原因分析
警告提示normal_loads组件对du_dt的导数为0,但从代码逻辑看,du_dt直接作为纵向加速度a_x传入载荷计算函数,理论上导数不应为0,可能的原因包括:
- JAX自动微分的数值近似问题:
smooth_max/smooth_min的mu值过小(当前为0.01),在初始值(du_dt=10)附近可能导致导数被数值近似为0。虽然du_dt=10远小于截断上限40,但极小的mu可能让JAX的自动微分出现精度损失。 - 载荷计算函数的导数异常:
calc_steady_state_normal_loads_jax函数内部可能存在对a_x导数为0的情况,比如某些条件分支或数值处理逻辑导致导数被抵消。 - OpenMDAO导数检查的阈值判定:OpenMDAO的导数检查可能将极小的非零导数判定为0,属于误报。
解决方案
- 单独验证载荷函数的导数:用JAX原生工具(如
jax.jacfwd)直接计算calc_steady_state_normal_loads_jax对a_x的导数,确认是否真的为0:import jax from cp_dynamics.load_transfer import calc_steady_state_normal_loads_jax # 构造测试输入 test_vehicle = ... # 你的VehicleConfig实例 test_u = jnp.array([10.0]) test_a_x = jnp.array([10.0]) test_a_y = jnp.array([1.0]) # 计算输出对a_x的导数 jac = jax.jacfwd(lambda ax: calc_steady_state_normal_loads_jax(test_vehicle, test_u, ax, test_a_y)[0])(test_a_x) print(jac) # 若输出非零,则问题不在载荷函数 - 调整smooth函数的mu值:将
mu从0.01调大到0.1或1.0,避免数值精度问题导致的导数近似为0。 - 显式声明组件偏导数:在
NormalODE的setup中添加setup_partials方法,显式声明输出对输入的依赖,帮助OpenMDAO更准确识别导数关系:def setup_partials(self): self.declare_partials(['fz_fr', 'fz_fl', 'fz_rr', 'fz_rl'], ['du_dt', 'r', 'u']) - 验证OpenMDAO组件的导数:调用
check_partials方法,查看normal_loads组件的实际导数数值:
若输出显示导数非零,说明警告是误报,可忽略并继续优化;若导数确实为0,需排查载荷计算函数的逻辑。prob = om.Problem() prob.model.add_subsystem('normal', NormalODE(num_nodes=1, vehicle=your_vehicle)) prob.setup() prob.run_model() prob.check_partials(compact_print=True) - 检查求解器配置:尝试将
cyc_group.linear_solver.options['assemble_jac']设为False,或调整Newton求解器的迭代参数,看是否影响导数计算。
内容的提问来源于stack exchange,提问作者Mickey Heine
相关产品推荐
相关产品推荐

