diff --git a/crazyflow/sim/integration.py b/crazyflow/sim/integration.py index e12510a8..ac6c2551 100644 --- a/crazyflow/sim/integration.py +++ b/crazyflow/sim/integration.py @@ -1,11 +1,4 @@ -"""Numerical integrators for the simulation dynamics. - -Note: - State limits are enforced by clipping after the integration (see e.g. - [clip_rotor_vel][crazyflow.sim.sim.clip_rotor_vel]). The derivatives are unclipped. Higher order - integration methods like RK4 evaluate the dynamics at intermediate states, so derivatives seen - by the model can transiently exceed the limits within a single step. -""" +"""Numerical integrators for the simulation dynamics.""" from enum import StrEnum from functools import partial diff --git a/crazyflow/sim/sim.py b/crazyflow/sim/sim.py index 77cad553..0baa5d86 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -158,16 +158,16 @@ def __init__( # simulation pipeline. for name, fn in build_control_fns(self.control, self.dynamics): append_fn(self.step_pipeline, fn, name=name) + # Keep the rotor command (RPM or thrust, see ``rotor_vel_limits``) within the motor limits + lower, upper = rotor_vel_limits(self.dynamics, self.drone) + clip_fn = partial(clip_rotor_vel_cmd, lower=lower, upper=upper, dynamics=self.dynamics) + append_fn(self.step_pipeline, clip_fn, name="clip_rotor_vel_cmd") integrate_fn = select_integrate_fn(self.integrator, select_dynamics_fn(self.dynamics)) append_fn(self.step_pipeline, integrate_fn, name="integration") - # Keep the rotor state (RPM or thrust, see ``rotor_vel_limits``) within its physical limits - lower, upper = rotor_vel_limits(self.dynamics, self.drone) - clip_fn = partial(clip_rotor_vel, lower=lower, upper=upper) - append_fn(self.step_pipeline, clip_fn, name="clip_rotor_vel") - append_fn(self.step_pipeline, increment_steps) # We never drop below -0.001 (drones can't pass through the floor). We use -0.001 to # enable checks for negative z sign append_fn(self.step_pipeline, clip_floor_pos) + append_fn(self.step_pipeline, increment_steps) self._reset = self.build_reset_fn() self._step = self.build_step_fn() @@ -718,10 +718,16 @@ def rotor_vel_limits(dynamics: Dynamics, drone: str) -> tuple[float, float]: return 4 * thrust_min, 4 * thrust_max -def clip_rotor_vel(data: SimData, lower: Array | float, upper: Array | float) -> SimData: - """Clip ``rotor_vel`` to ``[lower, upper]``.""" - rotor_vel = jnp.clip(data.states.rotor_vel, lower, upper) - return data.replace(states=data.states.replace(rotor_vel=rotor_vel)) +def clip_rotor_vel_cmd( + data: SimData, lower: Array | float, upper: Array | float, dynamics: Dynamics +) -> SimData: + """Clip the rotor command (RPM for first principles, collective thrust otherwise).""" + if dynamics == Dynamics.first_principles: + rotor_vel = jnp.clip(data.controls.rotor_vel, lower, upper) + return data.replace(controls=data.controls.replace(rotor_vel=rotor_vel)) + attitude = data.controls.attitude + cmd = attitude.cmd.at[..., -1].set(jnp.clip(attitude.cmd[..., -1], lower, upper)) + return data.replace(controls=data.controls.replace(attitude=attitude.replace(cmd=cmd))) @partial(jax.jit, static_argnames="device") diff --git a/docs/examples/index.md b/docs/examples/index.md index 20036a77..fc6a7913 100644 --- a/docs/examples/index.md +++ b/docs/examples/index.md @@ -62,9 +62,9 @@ Because the simulator is built entirely from JAX operations, `jax.grad` can diff --- -## Gradients and state clipping +## Gradients and command clipping -Hard state clips such as the `clip_rotor_vel` stage zero the gradients while the state is saturated. This example ramps a motor command beyond the rotor limits and compares the rotor state and the gradient of the vertical acceleration w.r.t. the command for three options: the default clip, a straight-through clip (clipped forward pass, unclipped gradients), and no clip. Replacing the default clip with the straight-through variant can help gradient-based methods such as trajectory optimization or policy learning, which would otherwise receive zero gradients whenever the motors saturate. +Hard clips such as the `clip_rotor_vel_cmd` stage zero the gradients while the command is saturated. This example ramps a motor command beyond the rotor limits and compares the rotor state and the gradient of the vertical acceleration w.r.t. the command for three options: the default clip, a straight-through clip (clipped forward pass, unclipped gradients), and no clip. Replacing the default clip with the straight-through variant can help gradient-based methods such as trajectory optimization or policy learning, which would otherwise receive zero gradients whenever the motors saturate. ```{ .python notest } diff --git a/docs/user-guide/oo-api.md b/docs/user-guide/oo-api.md index 441022d6..648e5846 100644 --- a/docs/user-guide/oo-api.md +++ b/docs/user-guide/oo-api.md @@ -44,6 +44,9 @@ Key constructor arguments: See [`Sim`][crazyflow.sim.Sim] in the API reference for the defaults and the full argument list. +!!! warning "Low simulation frequencies" + Explicit integrators are only stable if the step is small compared to the fastest time constant of the dynamics. For the drone models this is the rotor dynamics with time constants down to 20 ms. Euler overshoots once the step exceeds the time constant and diverges beyond twice the time constant. Keep `freq` well above 100 Hz. + ## Control methods All control methods take an array of shape `(n_worlds, n_drones, command_dim)` and stage it for the next `step` call. diff --git a/docs/user-guide/pipelines.md b/docs/user-guide/pipelines.md index 0fd3e3d3..9eb073c4 100644 --- a/docs/user-guide/pipelines.md +++ b/docs/user-guide/pipelines.md @@ -22,19 +22,22 @@ Both pipelines are constructed at `Sim` initialisation and compiled into a singl `sim.step_pipeline` contains multiple stages by default: 1. **Control functions** — convert the staged command through the control hierarchy (state → attitude → force/torque → rotor velocities, depending on the selected mode) -2. **Integrator** (`integration`) — advance the ODE one dynamics step (Euler, RK4, or symplectic Euler) -3. **Rotor clip** (`clip_rotor_vel`) — clip the motor speeds (first principles) or the collective thrust (fitted models) to the physical limits of the motors -4. **Step counter** (`increment_steps`) — increment `data.core.steps` -5. **Floor clip** (`clip_floor_pos`) — prevent drones from passing through the floor +2. **Rotor command clip** (`clip_rotor_vel_cmd`) — clip the commanded motor speeds (first principles) or collective thrust (so_rpy models) to the physical limits of the motors +3. **Integrator** (`integration`) — advance the ODE one dynamics step (Euler, RK4, or symplectic Euler) +4. **Floor clip** (`clip_floor_pos`) — prevent drones from passing through the floor +5. **Step counter** (`increment_steps`) — increment `data.core.steps` ```pycon >>> from crazyflow.sim import Sim >>> sim = Sim() >>> print(tuple(sim.step_pipeline.keys())) -('attitude_controller', 'force_torque_controller', 'integration', 'clip_rotor_vel', 'increment_steps', 'clip_floor_pos') +('attitude_controller', 'force_torque_controller', 'clip_rotor_vel_cmd', 'integration', 'clip_floor_pos', 'increment_steps') ``` +!!! note + We only clip the command, not the rotor state. Real motors deviate from their nominal thrust curve, and randomized motor parameters model exactly that, so the state has to be free to exceed the nominal limits. The command is always saturated, as it is on the real drone. + ## The reset pipeline `sim.reset_pipeline` holds a single `reset` stage that restores `SimData` to the default state. Every stage appended after it runs in order on the restored data. Each reset stage has the signature `(data: SimData, default_data: SimData, mask: Array | None) -> SimData`. The `default_data` argument holds the freshly-restored default state, which is useful for selectively reverting fields. @@ -131,9 +134,6 @@ remove_fn(sim.step_pipeline, "clip_floor_pos") sim.build_step_fn() ``` -!!! note - The rotor limits are enforced by the `clip_rotor_vel` stage, which clips the rotor state after each integration step. Higher order integration methods like RK4 evaluate the dynamics at intermediate, unclipped states, so the `rotor_vel` seen by the model can transiently exceed the limits within a single step. - ## Writing a custom stage A step pipeline function must have the signature `(SimData) -> SimData`. A reset pipeline function must have the signature `(SimData, SimData, Array | None) -> SimData` where the second argument is the default (freshly-restored) data. Both must be pure JAX functions with no Python-level side effects, so they can be traced and compiled. diff --git a/docs/user-guide/sim-overview.md b/docs/user-guide/sim-overview.md index 5100ac74..33ef4520 100644 --- a/docs/user-guide/sim-overview.md +++ b/docs/user-guide/sim-overview.md @@ -83,7 +83,7 @@ sim.step(sim.freq // sim.control_freq) # 500 // 100 = 5 dynamics steps, control ## The step and reset pipelines -Each call to `sim.step()` runs `sim.step_pipeline`, an ordered collection of named, pure JAX functions that transform `SimData`. By default it contains the control conversion functions, the numerical integrator, a clip of the rotor state to the physical motor limits, a step counter, and a floor clip. Similarly, `sim.reset_pipeline` is applied during `sim.reset()` and is empty by default. +Each call to `sim.step()` runs `sim.step_pipeline`, an ordered collection of named, pure JAX functions that transform `SimData`. By default it contains the control conversion functions, a clip of the rotor command to the physical motor limits, the numerical integrator, a floor clip, and a step counter. Similarly, `sim.reset_pipeline` is applied during `sim.reset()` and is empty by default. Both pipelines can be extended with custom functions for disturbances, domain randomization, or logging without modifying the core simulator. diff --git a/examples/jax/gradient_clipping.py b/examples/jax/gradient_clipping.py index 72c536c5..98c8f1e4 100644 --- a/examples/jax/gradient_clipping.py +++ b/examples/jax/gradient_clipping.py @@ -13,12 +13,12 @@ from crazyflow.sim.sim import rotor_vel_limits -def clip_rotor_vel_nonblocking(data: SimData, lower: float, upper: float) -> SimData: +def clip_rotor_vel_cmd_nonblocking(data: SimData, lower: float, upper: float) -> SimData: # Straight-through estimator: x + stop_gradient(clip(x) - x) evaluates to clip(x) in the # forward pass, while its derivative w.r.t. x is 1 in the backward pass - rotor_vel = data.states.rotor_vel - rotor_vel = rotor_vel + jax.lax.stop_gradient(jnp.clip(rotor_vel, lower, upper) - rotor_vel) - return data.replace(states=data.states.replace(rotor_vel=rotor_vel)) + cmd = data.controls.rotor_vel + cmd = cmd + jax.lax.stop_gradient(jnp.clip(cmd, lower, upper) - cmd) + return data.replace(controls=data.controls.replace(rotor_vel=cmd)) def rollout(sim: Sim, cmds: NDArray) -> tuple[NDArray, NDArray]: @@ -59,17 +59,17 @@ def main(plot: bool = False): cmds = np.concatenate([np.linspace(0, ramp, n), np.full(n, ramp), np.linspace(ramp, 0, n)]) # Option 1: keep the clipping as is. The rotor state respects the limits, but the gradient is - # zero while the state is saturated + # zero while the command is saturated results = {"clip (default)": rollout(sim, cmds)} - # Option 2: clip the state in the forward pass, but keep the gradients flowing in the backward - # pass (straight-through estimator) - clip_fn = partial(clip_rotor_vel_nonblocking, lower=lower, upper=upper) - replace_fn(sim.step_pipeline, clip_fn, "clip_rotor_vel") + # Option 2: clip the command in the forward pass, but keep the gradients flowing in the + # backward pass (straight-through estimator) + clip_fn = partial(clip_rotor_vel_cmd_nonblocking, lower=lower, upper=upper) + replace_fn(sim.step_pipeline, clip_fn, "clip_rotor_vel_cmd") results["nonblocking clip"] = rollout(sim, cmds) # Option 3: remove the clipping. Gradients always flow, but the state can leave the limits - remove_fn(sim.step_pipeline, "clip_rotor_vel") + remove_fn(sim.step_pipeline, "clip_rotor_vel_cmd") results["no clip"] = rollout(sim, cmds) sim.close() @@ -106,7 +106,7 @@ def plot_results( ax.set_xlabel("Time (s)") ax.legend() ax.grid(True) - fig.suptitle("Rotor state clipping and its effect on gradients") + fig.suptitle("Rotor command clipping and its effect on gradients") plt.tight_layout() plt.show() diff --git a/tests/unit/test_sim.py b/tests/unit/test_sim.py index 088e1fa1..ea6b69aa 100644 --- a/tests/unit/test_sim.py +++ b/tests/unit/test_sim.py @@ -450,7 +450,7 @@ def test_floor_penetration(dynamics: Dynamics): @pytest.mark.unit @pytest.mark.parametrize("integrator", Integrator) def test_rotor_vel_clip(integrator: Integrator): - """Test that the first-principles rotor state saturates at its physical limits.""" + """Test that the first principles rotor command is clipped to the motor limits.""" sim = Sim( dynamics=Dynamics.first_principles, control=Control.rotor_vel, @@ -460,14 +460,16 @@ def test_rotor_vel_clip(integrator: Integrator): lower, upper = rotor_vel_limits(Dynamics.first_principles, sim.drone) assert 0.0 < lower < upper - # States outside the limits are clipped after the integration for value, target in ((2 * upper, upper), (-upper, lower)): - states = sim.data.states.replace(rotor_vel=jnp.full_like(sim.data.states.rotor_vel, value)) + sim.reset() sim.rotor_vel_control(np.full((1, 1, 4), target)) - sim.data = sim.data.replace(states=states) sim.step() - assert jnp.all(sim.data.states.rotor_vel >= lower) - assert jnp.all(sim.data.states.rotor_vel <= upper) + rotor_vel_ref = sim.data.states.rotor_vel + sim.reset() + sim.rotor_vel_control(np.full((1, 1, 4), value)) + sim.step() + assert jnp.all(sim.data.controls.rotor_vel == target) + assert jnp.allclose(sim.data.states.rotor_vel, rotor_vel_ref) sim.close() @@ -477,18 +479,23 @@ def test_rotor_vel_clip(integrator: Integrator): ) @pytest.mark.parametrize("integrator", Integrator) def test_thrust_clip(dynamics: Dynamics, integrator: Integrator): - """Test that the thrust state of every so_rpy model is clipped to its physical limits.""" + """Test that the so_rpy thrust command is clipped to the motor limits.""" sim = Sim(dynamics=dynamics, control=Control.attitude, integrator=integrator, device="cpu") lower, upper = rotor_vel_limits(dynamics, sim.drone) assert 0.0 < lower < upper for value, target in ((2 * upper, upper), (-upper, lower)): - states = sim.data.states.replace(rotor_vel=jnp.full_like(sim.data.states.rotor_vel, value)) + # so_rpy has no thrust state and applies the command directly, so we compare the velocity + sim.reset() sim.attitude_control(np.array([[[0.0, 0.0, 0.0, target]]])) - sim.data = sim.data.replace(states=states) sim.step() - assert jnp.all(sim.data.states.rotor_vel >= lower) - assert jnp.all(sim.data.states.rotor_vel <= upper) + vel_ref, rotor_vel_ref = sim.data.states.vel, sim.data.states.rotor_vel + sim.reset() + sim.attitude_control(np.array([[[0.0, 0.0, 0.0, value]]])) + sim.step() + assert jnp.all(sim.data.controls.attitude.cmd[..., -1] == target) + assert jnp.allclose(sim.data.states.vel, vel_ref) + assert jnp.allclose(sim.data.states.rotor_vel, rotor_vel_ref) sim.close()