From 135afd11c31ce08cd816cb5927115d608c4eb2ca Mon Sep 17 00:00:00 2001 From: ratheron Date: Fri, 18 Sep 2026 15:23:49 +0200 Subject: [PATCH 1/2] Fix thrust clipping --- crazyflow/sim/sim.py | 23 +++++++++++++++++--- docs/examples/index.md | 2 +- docs/user-guide/pipelines.md | 15 ++++++------- examples/jax/gradient_clipping.py | 2 ++ tests/unit/test_sim.py | 35 ++++++++++++++++++++++++++----- 5 files changed, 61 insertions(+), 16 deletions(-) diff --git a/crazyflow/sim/sim.py b/crazyflow/sim/sim.py index 77cad553..97403565 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -158,16 +158,21 @@ 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 and state (RPM or thrust, see ``rotor_vel_limits``) within the + # physical limits. We need both: the command clip models the motor saturation and is the + # only limit for models without a rotor state (so_rpy), the state clip catches states set + # from outside the pipeline and integrator overshoot. + lower, upper = rotor_vel_limits(self.dynamics, self.drone) + clip_cmd_fn = partial(clip_rotor_vel_cmd, lower=lower, upper=upper, dynamics=self.dynamics) + append_fn(self.step_pipeline, clip_cmd_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() @@ -724,6 +729,18 @@ def clip_rotor_vel(data: SimData, lower: Array | float, upper: Array | float) -> 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") def seed_sim(data: SimData, seed: int, device: Device) -> SimData: """JIT-compiled seeding function.""" diff --git a/docs/examples/index.md b/docs/examples/index.md index 20036a77..c6a11d6d 100644 --- a/docs/examples/index.md +++ b/docs/examples/index.md @@ -64,7 +64,7 @@ Because the simulator is built entirely from JAX operations, `jax.grad` can diff ## Gradients and state 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 state clips such as the `clip_rotor_vel` stage zero the gradients while the state is saturated. This example removes the `clip_rotor_vel_cmd` stage, 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/pipelines.md b/docs/user-guide/pipelines.md index 0fd3e3d3..e7521e6d 100644 --- a/docs/user-guide/pipelines.md +++ b/docs/user-guide/pipelines.md @@ -22,19 +22,23 @@ 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` +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. **Rotor clip** (`clip_rotor_vel`) — clip the rotor state to the same limits 5. **Floor clip** (`clip_floor_pos`) — prevent drones from passing through the floor +6. **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_rotor_vel', 'clip_floor_pos', 'increment_steps') ``` +!!! note "Why two rotor clips?" + `clip_rotor_vel_cmd` models the motor saturation and is the only limit for models without a rotor state (`so_rpy`). `clip_rotor_vel` catches states set from outside the default pipeline, such as randomizations, and integrator overshoot. Since it runs after the integration, higher order integrators like RK4 can transiently exceed the limits within a step. + ## 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 +135,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/examples/jax/gradient_clipping.py b/examples/jax/gradient_clipping.py index 72c536c5..91c5fc6c 100644 --- a/examples/jax/gradient_clipping.py +++ b/examples/jax/gradient_clipping.py @@ -47,6 +47,8 @@ def acc_z(cmd: jax.Array, data: SimData) -> tuple[jax.Array, SimData]: def main(plot: bool = False): sim = Sim(control=Control.rotor_vel) lower, upper = rotor_vel_limits(sim.dynamics, sim.drone) + # Let the command reach the dynamics unclipped so that only the state clip affects the gradients + remove_fn(sim.step_pipeline, "clip_rotor_vel_cmd") # Start in the air so that the drone never reaches the floor, where the floor clipping would # zero the velocity and kill the gradients (see gradient.py) sim.data = sim.data.replace( diff --git a/tests/unit/test_sim.py b/tests/unit/test_sim.py index 088e1fa1..1df39d95 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 and rotor state are clipped to the limits.""" sim = Sim( dynamics=Dynamics.first_principles, control=Control.rotor_vel, @@ -460,11 +460,22 @@ 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)) + # Commands outside the limits are clipped before they enter the dynamics + sim.reset() sim.rotor_vel_control(np.full((1, 1, 4), target)) + sim.step() + 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) + # States outside the limits are clipped after the integration + sim.reset() + states = sim.data.states.replace(rotor_vel=jnp.full_like(sim.data.states.rotor_vel, value)) sim.data = sim.data.replace(states=states) + sim.rotor_vel_control(np.full((1, 1, 4), target)) sim.step() assert jnp.all(sim.data.states.rotor_vel >= lower) assert jnp.all(sim.data.states.rotor_vel <= upper) @@ -477,15 +488,29 @@ 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 and thrust state are clipped to the 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)) + # Commands outside the limits are clipped before they enter the dynamics. so_rpy has no + # thrust state and applies the command directly, so we compare the velocity as well + sim.reset() sim.attitude_control(np.array([[[0.0, 0.0, 0.0, target]]])) + sim.step() + 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) + # States outside the limits are clipped after the integration + sim.reset() + states = sim.data.states.replace(rotor_vel=jnp.full_like(sim.data.states.rotor_vel, value)) sim.data = sim.data.replace(states=states) + sim.attitude_control(np.array([[[0.0, 0.0, 0.0, target]]])) sim.step() assert jnp.all(sim.data.states.rotor_vel >= lower) assert jnp.all(sim.data.states.rotor_vel <= upper) From f3bc17305515424f213abda92b396e4c0031dc5c Mon Sep 17 00:00:00 2001 From: ratheron Date: Fri, 18 Sep 2026 17:01:50 +0200 Subject: [PATCH 2/2] Remove state clip --- crazyflow/sim/integration.py | 9 +-------- crazyflow/sim/sim.py | 17 +++-------------- docs/examples/index.md | 4 ++-- docs/user-guide/oo-api.md | 3 +++ docs/user-guide/pipelines.md | 11 +++++------ docs/user-guide/sim-overview.md | 2 +- examples/jax/gradient_clipping.py | 24 +++++++++++------------- tests/unit/test_sim.py | 24 +++--------------------- 8 files changed, 29 insertions(+), 65 deletions(-) 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 97403565..0baa5d86 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -158,17 +158,12 @@ 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 and state (RPM or thrust, see ``rotor_vel_limits``) within the - # physical limits. We need both: the command clip models the motor saturation and is the - # only limit for models without a rotor state (so_rpy), the state clip catches states set - # from outside the pipeline and integrator overshoot. + # 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_cmd_fn = partial(clip_rotor_vel_cmd, lower=lower, upper=upper, dynamics=self.dynamics) - append_fn(self.step_pipeline, clip_cmd_fn, name="clip_rotor_vel_cmd") + 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") - clip_fn = partial(clip_rotor_vel, lower=lower, upper=upper) - append_fn(self.step_pipeline, clip_fn, name="clip_rotor_vel") # 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) @@ -723,12 +718,6 @@ 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: diff --git a/docs/examples/index.md b/docs/examples/index.md index c6a11d6d..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 removes the `clip_rotor_vel_cmd` stage, 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 e7521e6d..9eb073c4 100644 --- a/docs/user-guide/pipelines.md +++ b/docs/user-guide/pipelines.md @@ -24,20 +24,19 @@ Both pipelines are constructed at `Sim` initialisation and compiled into a singl 1. **Control functions** — convert the staged command through the control hierarchy (state → attitude → force/torque → rotor velocities, depending on the selected mode) 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. **Rotor clip** (`clip_rotor_vel`) — clip the rotor state to the same limits -5. **Floor clip** (`clip_floor_pos`) — prevent drones from passing through the floor -6. **Step counter** (`increment_steps`) — increment `data.core.steps` +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', 'clip_rotor_vel_cmd', 'integration', 'clip_rotor_vel', 'clip_floor_pos', 'increment_steps') +('attitude_controller', 'force_torque_controller', 'clip_rotor_vel_cmd', 'integration', 'clip_floor_pos', 'increment_steps') ``` -!!! note "Why two rotor clips?" - `clip_rotor_vel_cmd` models the motor saturation and is the only limit for models without a rotor state (`so_rpy`). `clip_rotor_vel` catches states set from outside the default pipeline, such as randomizations, and integrator overshoot. Since it runs after the integration, higher order integrators like RK4 can transiently exceed the limits within a step. +!!! 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 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 91c5fc6c..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]: @@ -47,8 +47,6 @@ def acc_z(cmd: jax.Array, data: SimData) -> tuple[jax.Array, SimData]: def main(plot: bool = False): sim = Sim(control=Control.rotor_vel) lower, upper = rotor_vel_limits(sim.dynamics, sim.drone) - # Let the command reach the dynamics unclipped so that only the state clip affects the gradients - remove_fn(sim.step_pipeline, "clip_rotor_vel_cmd") # Start in the air so that the drone never reaches the floor, where the floor clipping would # zero the velocity and kill the gradients (see gradient.py) sim.data = sim.data.replace( @@ -61,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() @@ -108,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 1df39d95..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 command and rotor state are clipped to the limits.""" + """Test that the first principles rotor command is clipped to the motor limits.""" sim = Sim( dynamics=Dynamics.first_principles, control=Control.rotor_vel, @@ -461,7 +461,6 @@ def test_rotor_vel_clip(integrator: Integrator): assert 0.0 < lower < upper for value, target in ((2 * upper, upper), (-upper, lower)): - # Commands outside the limits are clipped before they enter the dynamics sim.reset() sim.rotor_vel_control(np.full((1, 1, 4), target)) sim.step() @@ -471,14 +470,6 @@ def test_rotor_vel_clip(integrator: Integrator): sim.step() assert jnp.all(sim.data.controls.rotor_vel == target) assert jnp.allclose(sim.data.states.rotor_vel, rotor_vel_ref) - # States outside the limits are clipped after the integration - sim.reset() - states = sim.data.states.replace(rotor_vel=jnp.full_like(sim.data.states.rotor_vel, value)) - sim.data = sim.data.replace(states=states) - sim.rotor_vel_control(np.full((1, 1, 4), target)) - sim.step() - assert jnp.all(sim.data.states.rotor_vel >= lower) - assert jnp.all(sim.data.states.rotor_vel <= upper) sim.close() @@ -488,14 +479,13 @@ def test_rotor_vel_clip(integrator: Integrator): ) @pytest.mark.parametrize("integrator", Integrator) def test_thrust_clip(dynamics: Dynamics, integrator: Integrator): - """Test that the so_rpy thrust command and thrust state are clipped to the 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)): - # Commands outside the limits are clipped before they enter the dynamics. so_rpy has no - # thrust state and applies the command directly, so we compare the velocity as well + # 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.step() @@ -506,14 +496,6 @@ def test_thrust_clip(dynamics: Dynamics, integrator: Integrator): 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) - # States outside the limits are clipped after the integration - sim.reset() - states = sim.data.states.replace(rotor_vel=jnp.full_like(sim.data.states.rotor_vel, value)) - sim.data = sim.data.replace(states=states) - sim.attitude_control(np.array([[[0.0, 0.0, 0.0, target]]])) - sim.step() - assert jnp.all(sim.data.states.rotor_vel >= lower) - assert jnp.all(sim.data.states.rotor_vel <= upper) sim.close()