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

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,可能的原因包括:

  1. JAX自动微分的数值近似问题:smooth_max/smooth_min的mu值过小(当前为0.01),在初始值(du_dt=10)附近可能导致导数被数值近似为0。虽然du_dt=10远小于截断上限40,但极小的mu可能让JAX的自动微分出现精度损失。
  2. 载荷计算函数的导数异常:calc_steady_state_normal_loads_jax函数内部可能存在对a_x导数为0的情况,比如某些条件分支或数值处理逻辑导致导数被抵消。
  3. OpenMDAO导数检查的阈值判定:OpenMDAO的导数检查可能将极小的非零导数判定为0,属于误报。

解决方案

  1. 单独验证载荷函数的导数:用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) # 若输出非零,则问题不在载荷函数
    
  2. 调整smooth函数的mu值:将mu从0.01调大到0.1或1.0,避免数值精度问题导致的导数近似为0。
  3. 显式声明组件偏导数:在NormalODE的setup中添加setup_partials方法,显式声明输出对输入的依赖,帮助OpenMDAO更准确识别导数关系:
    def setup_partials(self):
        self.declare_partials(['fz_fr', 'fz_fl', 'fz_rr', 'fz_rl'], ['du_dt', 'r', 'u'])
    
  4. 验证OpenMDAO组件的导数:调用check_partials方法,查看normal_loads组件的实际导数数值:
    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)
    
    若输出显示导数非零,说明警告是误报,可忽略并继续优化;若导数确实为0,需排查载荷计算函数的逻辑。
  5. 检查求解器配置:尝试将cyc_group.linear_solver.options['assemble_jac']设为False,或调整Newton求解器的迭代参数,看是否影响导数计算。

内容的提问来源于stack exchange,提问作者Mickey Heine

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 07:44:55