From 186c8bc7482c787b4caea8b532b105a8ba990ef1 Mon Sep 17 00:00:00 2001 From: ratheron Date: Tue, 6 Oct 2026 19:52:09 +0200 Subject: [PATCH 1/9] Expose states_deriv and fix symplectic integration --- crazyflow/sim/integration.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/crazyflow/sim/integration.py b/crazyflow/sim/integration.py index ac6c2551..d6beedda 100644 --- a/crazyflow/sim/integration.py +++ b/crazyflow/sim/integration.py @@ -86,7 +86,7 @@ def integrate(data: SimData, deriv: SimData, dt: float) -> SimData: states = states.replace( pos=next_pos, quat=next_quat, vel=next_vel, ang_vel=next_ang_vel, rotor_vel=next_rotor_vel ) - return data.replace(states=states) + return data.replace(states=states, states_deriv=states_deriv) def integrate_symplectic(data: SimData, deriv: SimData, dt: float) -> SimData: @@ -95,7 +95,7 @@ def integrate_symplectic(data: SimData, deriv: SimData, dt: float) -> SimData: 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 = states_deriv.acc, states_deriv.ang_acc, states_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 @@ -103,7 +103,7 @@ def integrate_symplectic(data: SimData, deriv: SimData, dt: float) -> SimData: states = states.replace( pos=next_pos, quat=next_quat, vel=next_vel, ang_vel=next_ang_vel, rotor_vel=next_rotor_vel ) - return data.replace(states=states) + return data.replace(states=states, states_deriv=states_deriv) @partial( From 9ee9f276740305d6e945bd718fb90e5b3ea4c164 Mon Sep 17 00:00:00 2001 From: ratheron Date: Wed, 7 Oct 2026 19:01:34 +0200 Subject: [PATCH 2/9] Remove `SimStateDeriv` from `SimData` --- .../dynamics/first_principles/dynamics.py | 9 +- crazyflow/dynamics/so_rpy/dynamics.py | 12 +- crazyflow/dynamics/so_rpy_rotor/dynamics.py | 9 +- .../dynamics/so_rpy_rotor_drag/dynamics.py | 9 +- crazyflow/sim/data.py | 2 - crazyflow/sim/integration.py | 47 ++++--- crazyflow/sim/sim.py | 5 +- docs/examples/index.md | 15 +++ docs/user-guide/sim-overview.md | 1 - examples/plugins/derivatives.py | 117 ++++++++++++++++++ tests/unit/test_utils.py | 1 - 11 files changed, 178 insertions(+), 49 deletions(-) create mode 100644 examples/plugins/derivatives.py diff --git a/crazyflow/dynamics/first_principles/dynamics.py b/crazyflow/dynamics/first_principles/dynamics.py index 38b85f22..bc274e23 100644 --- a/crazyflow/dynamics/first_principles/dynamics.py +++ b/crazyflow/dynamics/first_principles/dynamics.py @@ -33,7 +33,7 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData + from crazyflow.sim.data import SimData, SimStateDeriv @supports(rotor_dynamics=True) @@ -359,8 +359,10 @@ def create(drone: str, device: Device) -> Params: ) -def sim_dynamics(data: SimData) -> SimData: +def sim_dynamics(data: SimData) -> SimStateDeriv: """Compute the forces and torques from the first principle dynamics.""" + from crazyflow.sim.data import SimStateDeriv + params: Params = data.params vel, _, acc, ang_acc, rotor_acc = dynamics( pos=data.states.pos, @@ -373,7 +375,6 @@ def sim_dynamics(data: SimData) -> SimData: dist_t=data.states.torque, **params.__dict__, ) - states_deriv = data.states_deriv.replace( + return SimStateDeriv( 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/dynamics.py b/crazyflow/dynamics/so_rpy/dynamics.py index b031909f..01818b83 100644 --- a/crazyflow/dynamics/so_rpy/dynamics.py +++ b/crazyflow/dynamics/so_rpy/dynamics.py @@ -32,7 +32,7 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData + from crazyflow.sim.data import SimData, SimStateDeriv @supports(rotor_dynamics=False) @@ -360,8 +360,10 @@ def create(drone: str, device: Device) -> Params: ) -def sim_dynamics(data: SimData) -> SimData: +def sim_dynamics(data: SimData) -> SimStateDeriv: """Compute the forces and torques from the so_rpy dynamics.""" + from crazyflow.sim.data import SimStateDeriv + params: Params = data.params vel, _, acc, ang_acc = dynamics( pos=data.states.pos, @@ -373,7 +375,7 @@ def sim_dynamics(data: SimData) -> SimData: 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 = 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 ) - return data.replace(states_deriv=states_deriv) diff --git a/crazyflow/dynamics/so_rpy_rotor/dynamics.py b/crazyflow/dynamics/so_rpy_rotor/dynamics.py index 4bb3ad68..5193adce 100644 --- a/crazyflow/dynamics/so_rpy_rotor/dynamics.py +++ b/crazyflow/dynamics/so_rpy_rotor/dynamics.py @@ -34,7 +34,7 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData + from crazyflow.sim.data import SimData, SimStateDeriv @supports(rotor_dynamics=True) @@ -414,8 +414,10 @@ def create(drone: str, device: Device) -> Params: ) -def sim_dynamics(data: SimData) -> SimData: +def sim_dynamics(data: SimData) -> SimStateDeriv: """Compute the forces and torques from the so_rpy_rotor dynamics.""" + from crazyflow.sim.data import SimStateDeriv + params: Params = data.params vel, _, acc, ang_acc, rotor_acc = dynamics( pos=data.states.pos, @@ -428,7 +430,6 @@ def sim_dynamics(data: SimData) -> SimData: dist_t=data.states.torque, **params.__dict__, ) - states_deriv = data.states_deriv.replace( + return SimStateDeriv( 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/dynamics.py b/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py index 103ba270..329ecd16 100644 --- a/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py +++ b/crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py @@ -36,7 +36,7 @@ from jax import Device from crazyflow._typing import Array # To be changed to array_api_typing later - from crazyflow.sim.data import SimData + from crazyflow.sim.data import SimData, SimStateDeriv # Additional symbols specific to these dynamics @@ -451,8 +451,10 @@ def create(drone: str, device: Device) -> Params: ) -def sim_dynamics(data: SimData) -> SimData: +def sim_dynamics(data: SimData) -> SimStateDeriv: """Compute the forces and torques from the so_rpy_rotor_drag dynamics.""" + from crazyflow.sim.data import SimStateDeriv + params: Params = data.params vel, _, acc, ang_acc, rotor_acc = dynamics( pos=data.states.pos, @@ -465,7 +467,6 @@ def sim_dynamics(data: SimData) -> SimData: dist_t=data.states.torque, **params.__dict__, ) - states_deriv = data.states_deriv.replace( + return SimStateDeriv( 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/integration.py b/crazyflow/sim/integration.py index d6beedda..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 @@ -86,16 +83,16 @@ def integrate(data: SimData, deriv: SimData, dt: float) -> SimData: states = states.replace( pos=next_pos, quat=next_quat, vel=next_vel, ang_vel=next_ang_vel, rotor_vel=next_rotor_vel ) - return data.replace(states=states, states_deriv=states_deriv) + 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.acc, states_deriv.ang_acc, 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 @@ -103,7 +100,7 @@ def integrate_symplectic(data: SimData, deriv: SimData, dt: float) -> SimData: states = states.replace( pos=next_pos, quat=next_quat, vel=next_vel, ang_vel=next_ang_vel, rotor_vel=next_rotor_vel ) - return data.replace(states=states, states_deriv=states_deriv) + return data.replace(states=states) @partial( diff --git a/crazyflow/sim/sim.py b/crazyflow/sim/sim.py index 25ab886d..9737abe8 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -525,7 +525,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 +647,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 +663,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/examples/index.md b/docs/examples/index.md index 66dab31f..8bd19c53 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 to the step pipeline with plugins. The simulation does not store them by default because it costs a little performance and they are rarely needed. 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/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..802d5448 --- /dev/null +++ b/examples/plugins/derivatives.py @@ -0,0 +1,117 @@ +"""Example of adding the state derivatives to the step pipeline with plugins. + +Evaluating the dynamics gives the exact derivative at the current state. Finite differences give the +exact average derivative from the last to the current step. The simulation does not store either by +default because it costs performance and they are rarely needed. +""" + +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.dynamics.first_principles import sim_dynamics +from crazyflow.sim import Sim +from crazyflow.sim.data import SimStateDeriv +from crazyflow.sim.pipeline import append_fn + +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 at the current state.""" + return data.replace(plugins=data.plugins | {"states_deriv": sim_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): + sim = Sim(dynamics="first_principles", control="state", integrator="rk4") + 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) + + # Append after integration step + append_fn(sim.step_pipeline, 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(2 * 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) + for name in ("vel", "ang_vel", "acc", "ang_acc", "rotor_acc"): + x, x_fd = getattr(dynamics, name), getattr(fd, name) + print(f"{name}: max relative difference {np.abs(x - x_fd).max() / np.abs(x).max():.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 col, (name, title, unit) in enumerate(quantities): + x, x_fd = getattr(dynamics, name), getattr(fd, name) + x, x_fd = x[len(x) // 2 :], x_fd[len(x_fd) // 2 :] + t = np.arange(len(x)) / sim.control_freq + for i, axis in enumerate("xyz"): + label = f"{axis} finite differences" + axes[0, col].plot(t, x_fd[:, i], f"C{i}", lw=2.5, alpha=0.6, label=label) + axes[1, col].plot(t, x_fd[:, i] - x[:, i], f"C{i}", lw=0.8) + axes[0, col].plot(t, x, "k--", lw=0.8) + axes[0, col].set(title=title, ylabel=f"{title} [{unit}]") + axes[1, col].set(xlabel="Time [s]", ylabel=f"Finite differences - dynamics [{unit}]") + axes[0, 0].plot([], [], "k--", lw=0.8, label="dynamics") + 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 From 6eaf1c40bc400f44ddad1916b0944a0644e25512 Mon Sep 17 00:00:00 2001 From: ratheron Date: Wed, 7 Oct 2026 19:27:38 +0200 Subject: [PATCH 3/9] Remove function level imports by moving wrappers --- .../dynamics/first_principles/__init__.py | 9 +- .../dynamics/first_principles/dynamics.py | 22 ----- crazyflow/dynamics/so_rpy/__init__.py | 3 +- crazyflow/dynamics/so_rpy/dynamics.py | 22 ----- crazyflow/dynamics/so_rpy_rotor/__init__.py | 3 +- crazyflow/dynamics/so_rpy_rotor/dynamics.py | 22 ----- .../dynamics/so_rpy_rotor_drag/__init__.py | 3 +- .../dynamics/so_rpy_rotor_drag/dynamics.py | 22 ----- crazyflow/sim/dynamics.py | 89 +++++++++++++++++++ crazyflow/sim/sim.py | 10 ++- docs/api/index.md | 1 + examples/plugins/derivatives.py | 4 +- 12 files changed, 103 insertions(+), 107 deletions(-) create mode 100644 crazyflow/sim/dynamics.py 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 bc274e23..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, SimStateDeriv @supports(rotor_dynamics=True) @@ -357,24 +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) -> SimStateDeriv: - """Compute the forces and torques from the first principle dynamics.""" - from crazyflow.sim.data import SimStateDeriv - - 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__, - ) - return SimStateDeriv( - vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc - ) 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 01818b83..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, SimStateDeriv @supports(rotor_dynamics=False) @@ -358,24 +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) -> SimStateDeriv: - """Compute the forces and torques from the so_rpy dynamics.""" - from crazyflow.sim.data import SimStateDeriv - - 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__, - ) - 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 - ) 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 5193adce..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, SimStateDeriv @supports(rotor_dynamics=True) @@ -412,24 +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) -> SimStateDeriv: - """Compute the forces and torques from the so_rpy_rotor dynamics.""" - from crazyflow.sim.data import SimStateDeriv - - 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__, - ) - return SimStateDeriv( - vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc - ) 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 329ecd16..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, SimStateDeriv # Additional symbols specific to these dynamics @@ -449,24 +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) -> SimStateDeriv: - """Compute the forces and torques from the so_rpy_rotor_drag dynamics.""" - from crazyflow.sim.data import SimStateDeriv - - 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__, - ) - 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/dynamics.py b/crazyflow/sim/dynamics.py new file mode 100644 index 00000000..3ffd0665 --- /dev/null +++ b/crazyflow/sim/dynamics.py @@ -0,0 +1,89 @@ +"""Wrappers that wire `SimData` into the drone dynamics and return `SimStateDeriv`.""" + +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: + """Wire `SimData` into the first principles dynamics and return `SimStateDeriv`.""" + 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: + """Wire `SimData` into the so_rpy dynamics and return `SimStateDeriv`.""" + 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: + """Wire `SimData` into the so_rpy_rotor dynamics and return `SimStateDeriv`.""" + 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: + """Wire `SimData` into the so_rpy_rotor_drag dynamics and return `SimStateDeriv`.""" + 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/sim.py b/crazyflow/sim/sim.py index 9737abe8..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 diff --git a/docs/api/index.md b/docs/api/index.md index 8c602fab..932dd77a 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` | Wrappers that wire `SimData` into the drone dynamics and return `SimStateDeriv` | | `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/examples/plugins/derivatives.py b/examples/plugins/derivatives.py index 802d5448..56374ca1 100644 --- a/examples/plugins/derivatives.py +++ b/examples/plugins/derivatives.py @@ -17,9 +17,9 @@ import numpy as np from jax.scipy.spatial.transform import Rotation as R -from crazyflow.dynamics.first_principles import sim_dynamics 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 if TYPE_CHECKING: @@ -30,7 +30,7 @@ def dynamics_deriv(data: SimData) -> SimData: """Evaluate the dynamics at the current state.""" - return data.replace(plugins=data.plugins | {"states_deriv": sim_dynamics(data)}) + return data.replace(plugins=data.plugins | {"states_deriv": first_principles_dynamics(data)}) def finite_diff_deriv(data: SimData) -> SimData: From a759b499172c2056fac6a32c400ffbd46ab3ac37 Mon Sep 17 00:00:00 2001 From: ratheron Date: Wed, 7 Oct 2026 19:36:52 +0200 Subject: [PATCH 4/9] Fix docs --- docs/gen_ref_pages.py | 1 + 1 file changed, 1 insertion(+) 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) From 5e4e3e2a2530a09d6ebad077b83a6f0788247a4a Mon Sep 17 00:00:00 2001 From: Martin Schuck <57562633+amacati@users.noreply.github.com> Date: Thu, 8 Oct 2026 01:52:13 +0300 Subject: [PATCH 5/9] Update docstrings --- crazyflow/sim/dynamics.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/crazyflow/sim/dynamics.py b/crazyflow/sim/dynamics.py index 3ffd0665..04856868 100644 --- a/crazyflow/sim/dynamics.py +++ b/crazyflow/sim/dynamics.py @@ -1,4 +1,4 @@ -"""Wrappers that wire `SimData` into the drone dynamics and return `SimStateDeriv`.""" +"""Wrappers around the drone dynamics for `SimData` compatibility.""" from __future__ import annotations @@ -14,7 +14,7 @@ def first_principles_dynamics(data: SimData) -> SimStateDeriv: - """Wire `SimData` into the first principles dynamics and return `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, @@ -33,7 +33,7 @@ def first_principles_dynamics(data: SimData) -> SimStateDeriv: def so_rpy_dynamics(data: SimData) -> SimStateDeriv: - """Wire `SimData` into the so_rpy dynamics and return `SimStateDeriv`.""" + """Wrap the so_rpy dynamics.""" params: so_rpy.Params = data.params vel, _, acc, ang_acc = so_rpy.dynamics( pos=data.states.pos, @@ -52,7 +52,7 @@ def so_rpy_dynamics(data: SimData) -> SimStateDeriv: def so_rpy_rotor_dynamics(data: SimData) -> SimStateDeriv: - """Wire `SimData` into the so_rpy_rotor dynamics and return `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, @@ -71,7 +71,7 @@ def so_rpy_rotor_dynamics(data: SimData) -> SimStateDeriv: def so_rpy_rotor_drag_dynamics(data: SimData) -> SimStateDeriv: - """Wire `SimData` into the so_rpy_rotor_drag dynamics and return `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, From a4f9cbf655caf25490ce219ee320ff44b36c4646 Mon Sep 17 00:00:00 2001 From: Martin Schuck <57562633+amacati@users.noreply.github.com> Date: Thu, 8 Oct 2026 01:52:57 +0300 Subject: [PATCH 6/9] Update docs --- docs/api/index.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/api/index.md b/docs/api/index.md index 932dd77a..e1d991a3 100644 --- a/docs/api/index.md +++ b/docs/api/index.md @@ -9,7 +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` | Wrappers that wire `SimData` into the drone dynamics and return `SimStateDeriv` | +| `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 | From 35a7d6bad07ff3f2a4910fa5ecf121bfe2d3efa3 Mon Sep 17 00:00:00 2001 From: Martin Schuck <57562633+amacati@users.noreply.github.com> Date: Thu, 8 Oct 2026 01:54:45 +0300 Subject: [PATCH 7/9] Update docs --- docs/examples/index.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/examples/index.md b/docs/examples/index.md index 8bd19c53..fb613c3d 100644 --- a/docs/examples/index.md +++ b/docs/examples/index.md @@ -148,7 +148,7 @@ python examples/plugins/ground_effect.py ## State derivatives -Adding the state derivatives to the step pipeline with plugins. The simulation does not store them by default because it costs a little performance and they are rarely needed. 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. +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 } From fb13df62fc20aa91c5fb50c684e8fadf339670d3 Mon Sep 17 00:00:00 2001 From: ratheron Date: Thu, 8 Oct 2026 01:14:01 +0200 Subject: [PATCH 8/9] Fix example and add some comments --- examples/plugins/derivatives.py | 95 +++++++++++++++++---------------- 1 file changed, 50 insertions(+), 45 deletions(-) diff --git a/examples/plugins/derivatives.py b/examples/plugins/derivatives.py index 56374ca1..cfa63376 100644 --- a/examples/plugins/derivatives.py +++ b/examples/plugins/derivatives.py @@ -1,8 +1,10 @@ """Example of adding the state derivatives to the step pipeline with plugins. -Evaluating the dynamics gives the exact derivative at the current state. Finite differences give the -exact average derivative from the last to the current step. The simulation does not store either by -default because it costs performance and they are rarely needed. +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 @@ -20,7 +22,7 @@ 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 +from crazyflow.sim.pipeline import append_fn, insert_fn_before if TYPE_CHECKING: from crazyflow.sim.data import SimData @@ -29,7 +31,7 @@ def dynamics_deriv(data: SimData) -> SimData: - """Evaluate the dynamics at the current state.""" + """Evaluate the dynamics.""" return data.replace(plugins=data.plugins | {"states_deriv": first_principles_dynamics(data)}) @@ -57,34 +59,37 @@ def trajectory(t: float) -> np.ndarray: def main(plot: bool = True): - sim = Sim(dynamics="first_principles", control="state", integrator="rk4") - 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) - - # Append after integration step - append_fn(sim.step_pipeline, 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(2 * 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) - for name in ("vel", "ang_vel", "acc", "ang_acc", "rotor_acc"): - x, x_fd = getattr(dynamics, name), getattr(fd, name) - print(f"{name}: max relative difference {np.abs(x - x_fd).max() / np.abs(x).max():.1e}") + 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(2 * 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 @@ -94,22 +99,22 @@ def main(plot: bool = True): ("acc", "Linear acceleration", "m/s$^2$"), ("ang_acc", "Angular acceleration", "rad/s$^2$"), ) - for col, (name, title, unit) in enumerate(quantities): - x, x_fd = getattr(dynamics, name), getattr(fd, name) - x, x_fd = x[len(x) // 2 :], x_fd[len(x_fd) // 2 :] - t = np.arange(len(x)) / sim.control_freq - for i, axis in enumerate("xyz"): - label = f"{axis} finite differences" - axes[0, col].plot(t, x_fd[:, i], f"C{i}", lw=2.5, alpha=0.6, label=label) - axes[1, col].plot(t, x_fd[:, i] - x[:, i], f"C{i}", lw=0.8) - axes[0, col].plot(t, x, "k--", lw=0.8) - axes[0, col].set(title=title, ylabel=f"{title} [{unit}]") - axes[1, col].set(xlabel="Time [s]", ylabel=f"Finite differences - dynamics [{unit}]") - axes[0, 0].plot([], [], "k--", lw=0.8, label="dynamics") + 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) + x, x_fd = x[len(x) // 2 :], x_fd[len(x_fd) // 2 :] + 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.savefig("derivatives.png", dpi=300) plt.show() From 73a6c544e7de676b7c5e7f3934896c9da8ce9292 Mon Sep 17 00:00:00 2001 From: ratheron Date: Thu, 8 Oct 2026 12:54:25 +0200 Subject: [PATCH 9/9] Apply suggestions --- examples/plugins/derivatives.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/examples/plugins/derivatives.py b/examples/plugins/derivatives.py index cfa63376..e8d95fa7 100644 --- a/examples/plugins/derivatives.py +++ b/examples/plugins/derivatives.py @@ -77,7 +77,7 @@ def main(plot: bool = True): sim.build_step_fn() log = {"states_deriv": [], "fd_states_deriv": []} - for i in range(int(2 * DURATION * sim.control_freq)): + 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: @@ -102,7 +102,6 @@ def main(plot: bool = True): 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) - x, x_fd = x[len(x) // 2 :], x_fd[len(x_fd) // 2 :] 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) @@ -114,7 +113,6 @@ def main(plot: bool = True): for ax in axes.flat: ax.grid() fig.tight_layout() - plt.savefig("derivatives.png", dpi=300) plt.show()