diff --git a/crazyflow/dynamics/first_principles/__init__.py b/crazyflow/dynamics/first_principles/__init__.py index 62076062..25e180e9 100644 --- a/crazyflow/dynamics/first_principles/__init__.py +++ b/crazyflow/dynamics/first_principles/__init__.py @@ -81,11 +81,6 @@ moment of inertia, \(M\) is the \(3\times 4\) mixing matrix, and \(\mathbf{m}_z\) is its last row. """ -from crazyflow.dynamics.first_principles.dynamics import ( - Params, - dynamics, - sim_dynamics, - symbolic_dynamics, -) +from crazyflow.dynamics.first_principles.dynamics import Params, dynamics, symbolic_dynamics -__all__ = ["dynamics", "symbolic_dynamics", "sim_dynamics", "Params"] +__all__ = ["dynamics", "symbolic_dynamics", "Params"] diff --git a/crazyflow/dynamics/first_principles/dynamics.py b/crazyflow/dynamics/first_principles/dynamics.py index 38b85f22..899d8b03 100644 --- a/crazyflow/dynamics/first_principles/dynamics.py +++ b/crazyflow/dynamics/first_principles/dynamics.py @@ -33,7 +33,6 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData @supports(rotor_dynamics=True) @@ -357,23 +356,3 @@ def create(drone: str, device: Device) -> Params: drag_matrix=jnp.asarray(p["drag_matrix"], device=device), rotor_dyn_coef=jnp.asarray([p["rotor_dyn_coef"]], device=device), ) - - -def sim_dynamics(data: SimData) -> SimData: - """Compute the forces and torques from the first principle dynamics.""" - params: Params = data.params - vel, _, acc, ang_acc, rotor_acc = dynamics( - pos=data.states.pos, - quat=data.states.quat, - vel=data.states.vel, - ang_vel=data.states.ang_vel, - cmd=data.controls.rotor_vel, - rotor_vel=data.states.rotor_vel, - dist_f=data.states.force, - dist_t=data.states.torque, - **params.__dict__, - ) - states_deriv = data.states_deriv.replace( - vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc - ) - return data.replace(states_deriv=states_deriv) diff --git a/crazyflow/dynamics/so_rpy/__init__.py b/crazyflow/dynamics/so_rpy/__init__.py index c71d4324..a853e1c0 100644 --- a/crazyflow/dynamics/so_rpy/__init__.py +++ b/crazyflow/dynamics/so_rpy/__init__.py @@ -35,9 +35,8 @@ from crazyflow.dynamics.so_rpy.dynamics import ( Params, dynamics, - sim_dynamics, symbolic_dynamics, symbolic_dynamics_euler, ) -__all__ = ["Params", "dynamics", "sim_dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"] +__all__ = ["Params", "dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"] diff --git a/crazyflow/dynamics/so_rpy/dynamics.py b/crazyflow/dynamics/so_rpy/dynamics.py index b031909f..6503d8a3 100644 --- a/crazyflow/dynamics/so_rpy/dynamics.py +++ b/crazyflow/dynamics/so_rpy/dynamics.py @@ -32,7 +32,6 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData @supports(rotor_dynamics=False) @@ -358,22 +357,3 @@ def create(drone: str, device: Device) -> Params: rpy_rates_coef=jnp.asarray(p["rpy_rates_coef"], device=device), cmd_rpy_coef=jnp.asarray(p["cmd_rpy_coef"], device=device), ) - - -def sim_dynamics(data: SimData) -> SimData: - """Compute the forces and torques from the so_rpy dynamics.""" - params: Params = data.params - vel, _, acc, ang_acc = dynamics( - pos=data.states.pos, - quat=data.states.quat, - vel=data.states.vel, - ang_vel=data.states.ang_vel, - cmd=data.controls.attitude.cmd, - dist_f=data.states.force, - dist_t=data.states.torque, - **params.__dict__, - ) - states_deriv = data.states_deriv.replace( - vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc - ) - return data.replace(states_deriv=states_deriv) diff --git a/crazyflow/dynamics/so_rpy_rotor/__init__.py b/crazyflow/dynamics/so_rpy_rotor/__init__.py index 059b8fdc..4117c380 100644 --- a/crazyflow/dynamics/so_rpy_rotor/__init__.py +++ b/crazyflow/dynamics/so_rpy_rotor/__init__.py @@ -30,9 +30,8 @@ from crazyflow.dynamics.so_rpy_rotor.dynamics import ( Params, dynamics, - sim_dynamics, symbolic_dynamics, symbolic_dynamics_euler, ) -__all__ = ["Params", "dynamics", "sim_dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"] +__all__ = ["Params", "dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"] diff --git a/crazyflow/dynamics/so_rpy_rotor/dynamics.py b/crazyflow/dynamics/so_rpy_rotor/dynamics.py index 4bb3ad68..8a3222fa 100644 --- a/crazyflow/dynamics/so_rpy_rotor/dynamics.py +++ b/crazyflow/dynamics/so_rpy_rotor/dynamics.py @@ -34,7 +34,6 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData @supports(rotor_dynamics=True) @@ -412,23 +411,3 @@ def create(drone: str, device: Device) -> Params: rpy_rates_coef=jnp.asarray(p["rpy_rates_coef"], device=device), cmd_rpy_coef=jnp.asarray(p["cmd_rpy_coef"], device=device), ) - - -def sim_dynamics(data: SimData) -> SimData: - """Compute the forces and torques from the so_rpy_rotor dynamics.""" - params: Params = data.params - vel, _, acc, ang_acc, rotor_acc = dynamics( - pos=data.states.pos, - quat=data.states.quat, - vel=data.states.vel, - ang_vel=data.states.ang_vel, - rotor_vel=data.states.rotor_vel, - cmd=data.controls.attitude.cmd, - dist_f=data.states.force, - dist_t=data.states.torque, - **params.__dict__, - ) - states_deriv = data.states_deriv.replace( - vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc - ) - return data.replace(states_deriv=states_deriv) diff --git a/crazyflow/dynamics/so_rpy_rotor_drag/__init__.py b/crazyflow/dynamics/so_rpy_rotor_drag/__init__.py index 4113992e..92ed1842 100644 --- a/crazyflow/dynamics/so_rpy_rotor_drag/__init__.py +++ b/crazyflow/dynamics/so_rpy_rotor_drag/__init__.py @@ -32,9 +32,8 @@ from crazyflow.dynamics.so_rpy_rotor_drag.dynamics import ( Params, dynamics, - sim_dynamics, symbolic_dynamics, symbolic_dynamics_euler, ) -__all__ = ["Params", "dynamics", "sim_dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"] +__all__ = ["Params", "dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"] diff --git a/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py b/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py index 103ba270..552b8a34 100644 --- a/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py +++ b/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py @@ -36,7 +36,6 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData # Additional symbols specific to these dynamics @@ -449,23 +448,3 @@ def create(drone: str, device: Device) -> Params: cmd_rpy_coef=jnp.asarray(p["cmd_rpy_coef"], device=device), drag_matrix=jnp.asarray(p["drag_matrix"], device=device), ) - - -def sim_dynamics(data: SimData) -> SimData: - """Compute the forces and torques from the so_rpy_rotor_drag dynamics.""" - params: Params = data.params - vel, _, acc, ang_acc, rotor_acc = dynamics( - pos=data.states.pos, - quat=data.states.quat, - vel=data.states.vel, - ang_vel=data.states.ang_vel, - cmd=data.controls.attitude.cmd, - rotor_vel=data.states.rotor_vel, - dist_f=data.states.force, - dist_t=data.states.torque, - **params.__dict__, - ) - states_deriv = data.states_deriv.replace( - vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc - ) - return data.replace(states_deriv=states_deriv) diff --git a/crazyflow/sim/data.py b/crazyflow/sim/data.py index b66ec24e..b4d37878 100644 --- a/crazyflow/sim/data.py +++ b/crazyflow/sim/data.py @@ -277,8 +277,6 @@ def create( class SimData: states: SimState """State of the simulation.""" - states_deriv: SimStateDeriv - """Derivative of the state of the simulation.""" controls: SimControls """Drone controller data.""" params: SimParams diff --git a/crazyflow/sim/dynamics.py b/crazyflow/sim/dynamics.py new file mode 100644 index 00000000..04856868 --- /dev/null +++ b/crazyflow/sim/dynamics.py @@ -0,0 +1,89 @@ +"""Wrappers around the drone dynamics for `SimData` compatibility.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import jax.numpy as jnp + +from crazyflow.dynamics import first_principles, so_rpy, so_rpy_rotor, so_rpy_rotor_drag +from crazyflow.sim.data import SimStateDeriv + +if TYPE_CHECKING: + from crazyflow.sim.data import SimData + + +def first_principles_dynamics(data: SimData) -> SimStateDeriv: + """Wrap the first principles dynamics.""" + params: first_principles.Params = data.params + vel, _, acc, ang_acc, rotor_acc = first_principles.dynamics( + pos=data.states.pos, + quat=data.states.quat, + vel=data.states.vel, + ang_vel=data.states.ang_vel, + cmd=data.controls.rotor_vel, + rotor_vel=data.states.rotor_vel, + dist_f=data.states.force, + dist_t=data.states.torque, + **params.__dict__, + ) + return SimStateDeriv( + vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc + ) + + +def so_rpy_dynamics(data: SimData) -> SimStateDeriv: + """Wrap the so_rpy dynamics.""" + params: so_rpy.Params = data.params + vel, _, acc, ang_acc = so_rpy.dynamics( + pos=data.states.pos, + quat=data.states.quat, + vel=data.states.vel, + ang_vel=data.states.ang_vel, + cmd=data.controls.attitude.cmd, + dist_f=data.states.force, + dist_t=data.states.torque, + **params.__dict__, + ) + rotor_acc = jnp.zeros_like(data.states.rotor_vel) + return SimStateDeriv( + vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc + ) + + +def so_rpy_rotor_dynamics(data: SimData) -> SimStateDeriv: + """Wrap the so_rpy_rotor dynamics.""" + params: so_rpy_rotor.Params = data.params + vel, _, acc, ang_acc, rotor_acc = so_rpy_rotor.dynamics( + pos=data.states.pos, + quat=data.states.quat, + vel=data.states.vel, + ang_vel=data.states.ang_vel, + rotor_vel=data.states.rotor_vel, + cmd=data.controls.attitude.cmd, + dist_f=data.states.force, + dist_t=data.states.torque, + **params.__dict__, + ) + return SimStateDeriv( + vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc + ) + + +def so_rpy_rotor_drag_dynamics(data: SimData) -> SimStateDeriv: + """Wrap the so_rpy_rotor_drag dynamics.""" + params: so_rpy_rotor_drag.Params = data.params + vel, _, acc, ang_acc, rotor_acc = so_rpy_rotor_drag.dynamics( + pos=data.states.pos, + quat=data.states.quat, + vel=data.states.vel, + ang_vel=data.states.ang_vel, + cmd=data.controls.attitude.cmd, + rotor_vel=data.states.rotor_vel, + dist_f=data.states.force, + dist_t=data.states.torque, + **params.__dict__, + ) + return SimStateDeriv( + vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc + ) diff --git a/crazyflow/sim/integration.py b/crazyflow/sim/integration.py index ac6c2551..eb1e3519 100644 --- a/crazyflow/sim/integration.py +++ b/crazyflow/sim/integration.py @@ -10,7 +10,7 @@ from jax.numpy import vectorize from jax.scipy.spatial.transform import Rotation as R -from crazyflow.sim.data import SimData +from crazyflow.sim.data import SimData, SimStateDeriv class Integrator(StrEnum): @@ -20,7 +20,7 @@ class Integrator(StrEnum): default = euler -def euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData: +def euler(data: SimData, deriv_fn: Callable[[SimData], SimStateDeriv]) -> SimData: """Explicit Euler integration. Args: @@ -33,7 +33,7 @@ def euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData: return integrate(data, deriv_fn(data), dt=1 / data.core.freq) -def rk4(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData: +def rk4(data: SimData, deriv_fn: Callable[[SimData], SimStateDeriv]) -> SimData: """Runge-Kutta 4 integration. Args: @@ -44,14 +44,14 @@ def rk4(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData: The integrated simulation data structure. """ dt = 1 / data.core.freq - data_d1 = deriv_fn(data) - data_d2 = deriv_fn(integrate(data, data_d1, dt=dt / 2)) - data_d3 = deriv_fn(integrate(data, data_d2, dt=dt / 2)) - data_d4 = deriv_fn(integrate(data, data_d3, dt=dt)) - return integrate(data, rk4_average(data_d1, data_d2, data_d3, data_d4), dt=dt) + k1 = deriv_fn(data) + k2 = deriv_fn(integrate(data, k1, dt=dt / 2)) + k3 = deriv_fn(integrate(data, k2, dt=dt / 2)) + k4 = deriv_fn(integrate(data, k3, dt=dt)) + return integrate(data, rk4_average(k1, k2, k3, k4), dt=dt) -def symplectic_euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData: +def symplectic_euler(data: SimData, deriv_fn: Callable[[SimData], SimStateDeriv]) -> SimData: """Symplectic Euler integration. Args: @@ -61,24 +61,21 @@ def symplectic_euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> S return integrate_symplectic(data, deriv_fn(data), dt=1 / data.core.freq) -def rk4_average(k1: SimData, k2: SimData, k3: SimData, k4: SimData) -> SimData: +def rk4_average( + k1: SimStateDeriv, k2: SimStateDeriv, k3: SimStateDeriv, k4: SimStateDeriv +) -> SimStateDeriv: """Average four derivatives according to the RK4 rules.""" - data = k1 - k1, k2, k3, k4 = k1.states_deriv, k2.states_deriv, k3.states_deriv, k4.states_deriv - states_deriv = jax.tree.map( - lambda x1, x2, x3, x4: (x1 + 2 * x2 + 2 * x3 + x4) / 6, k1, k2, k3, k4 - ) - return data.replace(states_deriv=states_deriv) + return jax.tree.map(lambda x1, x2, x3, x4: (x1 + 2 * x2 + 2 * x3 + x4) / 6, k1, k2, k3, k4) -def integrate(data: SimData, deriv: SimData, dt: float) -> SimData: +def integrate(data: SimData, deriv: SimStateDeriv, dt: float) -> SimData: """Integrate the dynamics forward in time.""" - states, states_deriv = data.states, deriv.states_deriv + states = data.states pos, quat, vel, ang_vel = states.pos, states.quat, states.vel, states.ang_vel rotor_vel = states.rotor_vel - dpos, drot = states_deriv.vel, states_deriv.ang_vel - dvel, dang_vel, drotor_vel = states_deriv.acc, states_deriv.ang_acc, states_deriv.rotor_acc + dpos, drot = deriv.vel, deriv.ang_vel + dvel, dang_vel, drotor_vel = deriv.acc, deriv.ang_acc, deriv.rotor_acc next_pos, next_quat, next_vel, next_ang_vel, next_rotor_vel = _integrate( pos, quat, vel, ang_vel, rotor_vel, dpos, drot, dvel, dang_vel, drotor_vel, dt @@ -89,13 +86,13 @@ def integrate(data: SimData, deriv: SimData, dt: float) -> SimData: return data.replace(states=states) -def integrate_symplectic(data: SimData, deriv: SimData, dt: float) -> SimData: +def integrate_symplectic(data: SimData, deriv: SimStateDeriv, dt: float) -> SimData: """Integrate the dynamics forward in time.""" - states, states_deriv = data.states, deriv.states_deriv + states = data.states pos, quat, vel, ang_vel = states.pos, states.quat, states.vel, states.ang_vel rotor_vel = states.rotor_vel - dvel, dang_vel, drotor_vel = states_deriv.vel, states_deriv.ang_vel, states_deriv.rotor_acc + dvel, dang_vel, drotor_vel = deriv.acc, deriv.ang_acc, deriv.rotor_acc next_pos, next_quat, next_vel, next_ang_vel, next_rotor_vel = _integrate_symplectic( pos, quat, vel, ang_vel, rotor_vel, dvel, dang_vel, drotor_vel, dt diff --git a/crazyflow/sim/sim.py b/crazyflow/sim/sim.py index 25ab886d..ecd8c1c9 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -29,12 +29,14 @@ from crazyflow.drones import Drone from crazyflow.dynamics import Dynamics from crazyflow.dynamics import load_params as load_dynamics_params -from crazyflow.dynamics.first_principles import sim_dynamics as first_principles_dynamics -from crazyflow.dynamics.so_rpy import sim_dynamics as so_rpy_dynamics -from crazyflow.dynamics.so_rpy_rotor import sim_dynamics as so_rpy_rotor_dynamics -from crazyflow.dynamics.so_rpy_rotor_drag import sim_dynamics as so_rpy_rotor_drag_dynamics from crazyflow.exception import ConfigError, NotInitializedError from crazyflow.sim.data import SimControls, SimCore, SimData, SimParams, SimState, SimStateDeriv +from crazyflow.sim.dynamics import ( + first_principles_dynamics, + so_rpy_dynamics, + so_rpy_rotor_drag_dynamics, + so_rpy_rotor_dynamics, +) from crazyflow.sim.integration import Integrator, euler, rk4, symplectic_euler from crazyflow.sim.pipeline import append_fn from crazyflow.sim.sharding import WORLD_AXIS, build_sharded_data, build_sharded_mjx_data, placement @@ -525,7 +527,6 @@ def init_data( N, D = self.n_worlds, self.n_drones data = SimData( states=SimState.create(N, D, device), - states_deriv=SimStateDeriv.create(N, D, device), controls=SimControls.create( N, D, @@ -648,7 +649,7 @@ def build_control_fns( return stages -def select_dynamics_fn(dynamics: Dynamics) -> Callable[[SimData], SimData]: +def select_dynamics_fn(dynamics: Dynamics) -> Callable[[SimData], SimStateDeriv]: """Select the dynamics function for the given dynamics mode.""" match dynamics: case Dynamics.first_principles: @@ -664,7 +665,7 @@ def select_dynamics_fn(dynamics: Dynamics) -> Callable[[SimData], SimData]: def select_integrate_fn( - integrator: Integrator, dynamics_fn: Callable[[SimData], SimData] + integrator: Integrator, dynamics_fn: Callable[[SimData], SimStateDeriv] ) -> Callable[[SimData], SimData]: """Select the integration function for the given dynamics and integrator mode.""" match integrator: diff --git a/docs/api/index.md b/docs/api/index.md index 8c602fab..e1d991a3 100644 --- a/docs/api/index.md +++ b/docs/api/index.md @@ -9,6 +9,7 @@ This section is auto-generated from the Crazyflow source code using [mkdocstring | `crazyflow.sim` | Core `Sim` class and simulation pipeline | | `crazyflow.sim.pipeline` | `OrderedDict`-based pipeline helpers (`append_fn`, `insert_fn_before`, `replace_fn`, etc.) | | `crazyflow.sim.data` | `SimData`, `SimState`, `SimControls`, `SimParams`, `SimCore` pytrees | +| `crazyflow.sim.dynamics` | Wrap the drone dynamics for compatibility with `SimData` | | `crazyflow.sim.functional` | Pure functional control API for use inside `jax.jit` | | `crazyflow.sim.sharding` | Placement of the simulation data on multiple devices | | `crazyflow.dynamics` | `Dynamics` enum and dynamics implementations | diff --git a/docs/examples/index.md b/docs/examples/index.md index 66dab31f..fb613c3d 100644 --- a/docs/examples/index.md +++ b/docs/examples/index.md @@ -146,6 +146,21 @@ python examples/plugins/ground_effect.py --- +## State derivatives + +Adding the state derivatives, which are not stored by default, to the step pipeline with plugins. One plugin calls the dynamics at the end of each step, which gives the exact derivative at the current state. The other uses finite differences of the states, which give the exact average derivative from the last to the current step. Both store a `SimStateDeriv` in the plugins dict. The plot compares them while the drone flies a figure-eight. + + +```{ .python notest } +--8<-- "examples/plugins/derivatives.py" +``` + +```bash +python examples/plugins/derivatives.py +``` + +--- + ## Cameras and RGBD Offscreen rendering returns RGB-D images on every frame. The FPV camera (`fpv_cam`) is attached to the drone and moves with it. diff --git a/docs/gen_ref_pages.py b/docs/gen_ref_pages.py index 88e40cbf..f5365c25 100644 --- a/docs/gen_ref_pages.py +++ b/docs/gen_ref_pages.py @@ -47,6 +47,7 @@ * [drones](crazyflow/drones/index.md) * [sim](crazyflow/sim/index.md) * [sim.data](crazyflow/sim/data.md) + * [sim.dynamics](crazyflow/sim/dynamics.md) * [sim.functional](crazyflow/sim/functional.md) * [sim.integration](crazyflow/sim/integration.md) * [sim.pipeline](crazyflow/sim/pipeline.md) diff --git a/docs/user-guide/sim-overview.md b/docs/user-guide/sim-overview.md index 33ef4520..326ea361 100644 --- a/docs/user-guide/sim-overview.md +++ b/docs/user-guide/sim-overview.md @@ -25,7 +25,6 @@ All simulation state is stored in `sim.data`, a `SimData` pytree. The main sub-t | Field | Type | Description | |---|---|---| | `states` | `SimState` | Current kinematic state of every drone | -| `states_deriv` | `SimStateDeriv` | Time derivatives computed by the dynamics | | `controls` | `SimControls` | Staged commands and controller state | | `params` | `SimParams` | Physical parameters (mass, inertia, motor constants, …) | | `core` | `SimCore` | Metadata: step count, frequency, RNG key, device | diff --git a/examples/plugins/derivatives.py b/examples/plugins/derivatives.py new file mode 100644 index 00000000..e8d95fa7 --- /dev/null +++ b/examples/plugins/derivatives.py @@ -0,0 +1,120 @@ +"""Example of adding the state derivatives to the step pipeline with plugins. + +We compare two methods for the state derivatives: evaluating the dynamics and finite differences. +The dynamics need the current state and input, so we evaluate them before integration. At the end of +the step, they are the derivative at the previous state. Finite differences are taken after +integration and give the derivative from the previous to the current state. For Euler integration, +both methods are identical up to numerical noise. Other integrators differ, as shown here for RK4. +""" + +from __future__ import annotations + +import os +from typing import TYPE_CHECKING + +os.environ["SCIPY_ARRAY_API"] = "1" + +import jax +import jax.numpy as jnp +import numpy as np +from jax.scipy.spatial.transform import Rotation as R + +from crazyflow.sim import Sim +from crazyflow.sim.data import SimStateDeriv +from crazyflow.sim.dynamics import first_principles_dynamics +from crazyflow.sim.pipeline import append_fn, insert_fn_before + +if TYPE_CHECKING: + from crazyflow.sim.data import SimData + +DURATION = 10.0 # s, one loop of the figure-eight + + +def dynamics_deriv(data: SimData) -> SimData: + """Evaluate the dynamics.""" + return data.replace(plugins=data.plugins | {"states_deriv": first_principles_dynamics(data)}) + + +def finite_diff_deriv(data: SimData) -> SimData: + """Differentiate the states over the last step.""" + prev, states, freq = data.plugins["prev_states"], data.states, data.core.freq + rot = R.from_quat(prev.quat).inv() * R.from_quat(states.quat) + deriv = SimStateDeriv( + vel=(states.pos - prev.pos) * freq, + ang_vel=rot.as_rotvec() * freq, + acc=(states.vel - prev.vel) * freq, + ang_acc=(states.ang_vel - prev.ang_vel) * freq, + rotor_acc=(states.rotor_vel - prev.rotor_vel) * freq, + ) + return data.replace(plugins=data.plugins | {"fd_states_deriv": deriv, "prev_states": states}) + + +def trajectory(t: float) -> np.ndarray: + """Return a figure-eight state command.""" + omega = 2 * np.pi / DURATION + cmd = np.zeros((1, 1, 16)) + cmd[..., :3] = [2 * np.sin(omega * t), np.sin(2 * omega * t), 1.0] + cmd[..., 9:13] = R.from_euler("z", 0.0).as_quat() + return cmd + + +def main(plot: bool = True): + results = {} + for integrator in ("euler", "rk4"): + sim = Sim(dynamics="first_principles", control="state", integrator=integrator) + pos = jnp.asarray(trajectory(0.0)[..., :3], device=sim.device) + sim.data = sim.data.replace(states=sim.data.states.replace(pos=pos)) + plugins = { + "states_deriv": SimStateDeriv.create(sim.n_worlds, sim.n_drones, sim.device), + "fd_states_deriv": SimStateDeriv.create(sim.n_worlds, sim.n_drones, sim.device), + "prev_states": sim.data.states, + } + sim.data = sim.data.replace(plugins=sim.data.plugins | plugins) + + insert_fn_before(sim.step_pipeline, "integration", dynamics_deriv) + append_fn(sim.step_pipeline, finite_diff_deriv) + sim.build_default_data() + sim.build_step_fn() + + log = {"states_deriv": [], "fd_states_deriv": []} + for i in range(int(DURATION * sim.control_freq)): + sim.state_control(trajectory(i / sim.control_freq)) + sim.step(sim.freq // sim.control_freq) + for key in log: + log[key].append(jax.tree.map(lambda x: np.asarray(x[0, 0]), sim.data.plugins[key])) + sim.close() + + dynamics, fd = (jax.tree.map(lambda *x: np.stack(x), *log[key]) for key in log) + results[integrator] = dynamics, fd + for name in ("vel", "ang_vel", "acc", "ang_acc", "rotor_acc"): + x, x_fd = getattr(dynamics, name), getattr(fd, name) + diff = np.abs(x - x_fd).max() / np.abs(x).max() + print(f"{integrator} {name}: max relative difference {diff:.1e}") + + if plot: + import matplotlib.pyplot as plt + + fig, axes = plt.subplots(2, 2, sharex=True, figsize=(12, 7)) + quantities = ( + ("acc", "Linear acceleration", "m/s$^2$"), + ("ang_acc", "Angular acceleration", "rad/s$^2$"), + ) + for row, (integrator, (dynamics, fd)) in enumerate(results.items()): + for col, (name, title, unit) in enumerate(quantities): + x, x_fd = getattr(dynamics, name), getattr(fd, name) + t = np.arange(len(x)) / sim.control_freq + for i, axis in enumerate("xyz"): + axes[row, col].plot(t, x_fd[:, i] - x[:, i], f"C{i}", lw=0.8, label=axis) + ylabel = f"Finite differences - dynamics [{unit}]" + axes[row, col].set(title=f"{title} ({integrator})", ylabel=ylabel) + for ax in axes[-1]: + ax.set(xlabel="Time [s]") + axes[0, 0].legend() + for ax in axes.flat: + ax.grid() + fig.tight_layout() + plt.show() + + +if __name__ == "__main__": + main() diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index d85fd89a..e7fb9076 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -67,7 +67,6 @@ def test_world_mask(): # Check that the mask correctly handles cases that previously failed under shape-based detection mask = world_mask(Sim(n_worlds=3, control=Control.attitude).data) assert mask.states.pos - assert mask.states_deriv.acc assert mask.core.steps assert not mask.params.mass assert mask.controls.attitude.cmd