diff --git a/crazyflow/control/__init__.py b/crazyflow/control/__init__.py index 90ebc9e3..bb651c4b 100644 --- a/crazyflow/control/__init__.py +++ b/crazyflow/control/__init__.py @@ -14,7 +14,7 @@ __all__ = [] -from crazyflow.control.core import Control, parametrize +from crazyflow.control.core import Control, load_params, parametrize from crazyflow.control.mellinger import attitude2force_torque as mellinger_attitude2force_torque from crazyflow.control.mellinger import state2attitude as mellinger_state2attitude @@ -23,4 +23,4 @@ "mellinger_attitude2force_torque": mellinger_attitude2force_torque, } -__all__ = ["Control", "parametrize"] +__all__ = ["Control", "load_params", "parametrize"] diff --git a/crazyflow/dynamics/__init__.py b/crazyflow/dynamics/__init__.py index bd90eed6..3c450880 100644 --- a/crazyflow/dynamics/__init__.py +++ b/crazyflow/dynamics/__init__.py @@ -14,13 +14,13 @@ from typing import Callable -from crazyflow.dynamics.core import Dynamics, parametrize +from crazyflow.dynamics.core import Dynamics, load_params, parametrize from crazyflow.dynamics.first_principles import dynamics as _first_principles_dynamics from crazyflow.dynamics.so_rpy import dynamics as _so_rpy_dynamics from crazyflow.dynamics.so_rpy_rotor import dynamics as _so_rpy_rotor_dynamics from crazyflow.dynamics.so_rpy_rotor_drag import dynamics as _so_rpy_rotor_drag_dynamics -__all__ = ["parametrize", "available_dynamics", "dynamics_features", "Dynamics"] +__all__ = ["parametrize", "load_params", "available_dynamics", "dynamics_features", "Dynamics"] available_dynamics: dict[str, Callable] = { diff --git a/crazyflow/envs/drone_env.py b/crazyflow/envs/drone_env.py index c5ed6ab0..12930a52 100644 --- a/crazyflow/envs/drone_env.py +++ b/crazyflow/envs/drone_env.py @@ -12,8 +12,7 @@ from numpy.typing import NDArray from crazyflow.control import Control -from crazyflow.control.core import load_params -from crazyflow.control.mellinger import force_torque2rotor_vel +from crazyflow.drones import load_params from crazyflow.dynamics import Dynamics from crazyflow.sim import Sim from crazyflow.sim.data import SimData @@ -33,7 +32,7 @@ def action_space(control_type: Control, drone: str) -> spaces.Box: """ match control_type: case Control.attitude: - params = load_params(force_torque2rotor_vel, drone) + params = load_params(drone) thrust_min, thrust_max = params["thrust_min"] * 4, params["thrust_max"] * 4 return spaces.Box( np.array([-np.pi / 2, -np.pi / 2, -np.pi / 2, thrust_min], dtype=np.float32), diff --git a/docs/user-guide/control/parametrize.md b/docs/user-guide/control/parametrize.md index 4742a7ad..0795cea3 100644 --- a/docs/user-guide/control/parametrize.md +++ b/docs/user-guide/control/parametrize.md @@ -77,10 +77,10 @@ rpyt, _ = ctrl(pos, quat, vel, cmd) ## Loading raw parameters -Use [`load_params`][crazyflow.control.core.load_params] to inspect or override the values that `parametrize` would bind for a specific controller function: +Use [`load_params`][crazyflow.control.load_params] to inspect or override the values that `parametrize` would bind for a specific controller function: ```python -from crazyflow.control.core import load_params +from crazyflow.control import load_params from crazyflow.control.mellinger import state2attitude params = load_params(state2attitude, "cf2x_L250") diff --git a/docs/user-guide/dynamics/parametrize.md b/docs/user-guide/dynamics/parametrize.md index 59e549f9..1b7a2ecf 100644 --- a/docs/user-guide/dynamics/parametrize.md +++ b/docs/user-guide/dynamics/parametrize.md @@ -114,10 +114,10 @@ parametrized_dynamics = parametrize(dynamics, drone="cf2x_T350") ## Loading raw parameters -If you need the parameter values directly, for example, to pass them to [`symbolic_dynamics`](symbolic.md), use [`load_params`][crazyflow.dynamics.core.load_params]: +If you need the parameter values directly, for example, to pass them to [`symbolic_dynamics`](symbolic.md), use [`load_params`][crazyflow.dynamics.load_params]: ```python { .python continuation } -from crazyflow.dynamics.core import load_params +from crazyflow.dynamics import load_params params = load_params(dynamics, "cf2x_L250") params["mass"] # 0.0319 diff --git a/tests/integration/test_interfaces.py b/tests/integration/test_interfaces.py index a2cbc9dd..34b1036a 100644 --- a/tests/integration/test_interfaces.py +++ b/tests/integration/test_interfaces.py @@ -3,8 +3,7 @@ import pytest from scipy.spatial.transform import Rotation as R -from crazyflow.control import Control, parametrize -from crazyflow.control.core import load_params +from crazyflow.control import Control, load_params, parametrize from crazyflow.control.mellinger import force_torque2rotor_vel, state2attitude from crazyflow.control.transform import motor_force2rotor_vel from crazyflow.sim import Dynamics, Sim diff --git a/tests/unit/control/test_core.py b/tests/unit/control/test_core.py index c8e77437..1fd698cf 100644 --- a/tests/unit/control/test_core.py +++ b/tests/unit/control/test_core.py @@ -6,7 +6,7 @@ import array_api_strict import pytest -from crazyflow.control.core import load_params, parametrize +from crazyflow.control import load_params, parametrize from crazyflow.control.mellinger import ( attitude2force_torque, force_torque2rotor_vel, diff --git a/tests/unit/control/test_mellinger.py b/tests/unit/control/test_mellinger.py index 0876cc71..adb96f5f 100644 --- a/tests/unit/control/test_mellinger.py +++ b/tests/unit/control/test_mellinger.py @@ -5,8 +5,7 @@ import numpy as np import pytest -from crazyflow.control import parametrize -from crazyflow.control.core import load_params +from crazyflow.control import load_params, parametrize from crazyflow.control.mellinger import ( attitude2force_torque, force_torque2rotor_vel, diff --git a/tests/unit/dynamics/test_parametrization.py b/tests/unit/dynamics/test_parametrization.py index 1472a22d..e2f87c71 100644 --- a/tests/unit/dynamics/test_parametrization.py +++ b/tests/unit/dynamics/test_parametrization.py @@ -7,8 +7,15 @@ import pytest from crazyflow.drones import available_drones -from crazyflow.dynamics import available_dynamics -from crazyflow.dynamics.core import parametrize +from crazyflow.dynamics import available_dynamics, load_params, parametrize + + +@pytest.mark.unit +@pytest.mark.parametrize("dynamics_name, dynamics", available_dynamics.items()) +@pytest.mark.parametrize("drone", available_drones) +def test_dynamics_parameter_loading(dynamics_name: str, dynamics: Callable, drone: str) -> None: + """Check that parameters can be loaded for all available dynamics and drones.""" + load_params(dynamics, drone) @pytest.mark.unit