Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 1 addition & 8 deletions crazyflow/sim/integration.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down
24 changes: 15 additions & 9 deletions crazyflow/sim/sim.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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")
Expand Down
4 changes: 2 additions & 2 deletions docs/examples/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

<!-- notest: imported script, covered by tests/integration/test_examples.py -->
```{ .python notest }
Expand Down
3 changes: 3 additions & 0 deletions docs/user-guide/oo-api.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
16 changes: 8 additions & 8 deletions docs/user-guide/pipelines.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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.
Expand Down
2 changes: 1 addition & 1 deletion docs/user-guide/sim-overview.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
22 changes: 11 additions & 11 deletions examples/jax/gradient_clipping.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]:
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()

Expand Down
29 changes: 18 additions & 11 deletions tests/unit/test_sim.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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()


Expand All @@ -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()


Expand Down
Loading