From b83310fde81af8a1b51bc252360695ecff83d9c5 Mon Sep 17 00:00:00 2001 From: Yuguo Shan <237990095+Mikasa0503@users.noreply.github.com> Date: Sun, 20 Sep 2026 20:49:28 +0800 Subject: [PATCH 1/4] test: cover all simulation integrators --- tests/unit/test_sim.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/tests/unit/test_sim.py b/tests/unit/test_sim.py index ea6b69aa..405ce524 100644 --- a/tests/unit/test_sim.py +++ b/tests/unit/test_sim.py @@ -228,6 +228,21 @@ def test_sim_step(n_worlds: int, n_drones: int, dynamics: Dynamics, control: Con sim.step(2) +@pytest.mark.unit +@pytest.mark.parametrize("integrator", Integrator) +def test_sim_step_integrators(integrator: Integrator): + """Every integration scheme advances a default simulation without invalid state.""" + sim = Sim(integrator=integrator, device="cpu") + sim.step(2) + + assert jnp.all(sim.data.core.steps == 2) + assert jnp.all(jnp.isfinite(sim.data.states.pos)) + assert jnp.all(jnp.isfinite(sim.data.states.quat)) + assert jnp.all(jnp.isfinite(sim.data.states.vel)) + assert jnp.all(jnp.isfinite(sim.data.states.ang_vel)) + sim.close() + + @pytest.mark.unit def test_state_control_forwards_body_rates(): """State control must forward the body rates of the command to the attitude controller.""" From a6795fb0a9a1d3f5d50e7124919e04c3b8e5b677 Mon Sep 17 00:00:00 2001 From: Yuguo Shan <237990095+Mikasa0503@users.noreply.github.com> Date: Sun, 20 Sep 2026 21:56:26 +0800 Subject: [PATCH 2/4] test: validate all simulation state tensors --- tests/unit/test_sim.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/unit/test_sim.py b/tests/unit/test_sim.py index 405ce524..4ffa5316 100644 --- a/tests/unit/test_sim.py +++ b/tests/unit/test_sim.py @@ -240,6 +240,9 @@ def test_sim_step_integrators(integrator: Integrator): assert jnp.all(jnp.isfinite(sim.data.states.quat)) assert jnp.all(jnp.isfinite(sim.data.states.vel)) assert jnp.all(jnp.isfinite(sim.data.states.ang_vel)) + assert jnp.all(jnp.isfinite(sim.data.states.force)) + assert jnp.all(jnp.isfinite(sim.data.states.torque)) + assert jnp.all(jnp.isfinite(sim.data.states.rotor_vel)) sim.close() From 5ac9ed78351e9ebb39d7f576b3c968487968c723 Mon Sep 17 00:00:00 2001 From: Yuguo Shan <237990095+Mikasa0503@users.noreply.github.com> Date: Thu, 8 Oct 2026 12:34:14 +0800 Subject: [PATCH 3/4] test: compare full simulation data across integrators --- tests/unit/test_integrators.py | 54 ++++++++++++++++++++++++++++++++++ tests/unit/test_sim.py | 18 ------------ 2 files changed, 54 insertions(+), 18 deletions(-) create mode 100644 tests/unit/test_integrators.py diff --git a/tests/unit/test_integrators.py b/tests/unit/test_integrators.py new file mode 100644 index 00000000..5851f565 --- /dev/null +++ b/tests/unit/test_integrators.py @@ -0,0 +1,54 @@ +"""Unit tests for simulation integrators.""" + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from crazyflow.sim import Sim +from crazyflow.sim.integration import Integrator + + +@pytest.mark.unit +@pytest.mark.parametrize("integrator", Integrator) +def test_sim_step_integrators(integrator: Integrator, device: str): + """Every integrator stays close to the default after two simulation steps.""" + sim = Sim(integrator=integrator, device=device) + default_sim = Sim(device=device) + try: + sim.step(2) + default_sim.step(2) + + assert jnp.all(sim.data.core.steps == 2) + assert jnp.all(default_sim.data.core.steps == 2) + + data, structure = jax.tree.flatten_with_path(sim.data) + default_data, default_structure = jax.tree.flatten(default_sim.data) + assert structure == default_structure, "Simulation data structures must match" + for (path, value), default_value in zip(data, default_data, strict=True): + name = jax.tree_util.keystr(path) + if isinstance(value, jnp.ndarray): + assert type(value) is type(default_value), f"{name}: Type mismatch" + assert value.shape == default_value.shape, f"{name}: Shape mismatch" + assert value.dtype == default_value.dtype, f"{name}: Dtype mismatch" + assert value.device == default_value.device, f"{name}: Device mismatch" + if jax.dtypes.issubdtype(value.dtype, jax.dtypes.prng_key): + np.testing.assert_array_equal( + jax.random.key_data(value), jax.random.key_data(default_value), err_msg=name + ) + else: + assert jnp.all(jnp.isfinite(value)), f"{name}: Non-finite integrator data" + assert jnp.all(jnp.isfinite(default_value)), f"{name}: Non-finite default data" + if jnp.issubdtype(value.dtype, jnp.inexact): + # At 500 Hz, position differences are O(g * dt**2), below 1e-4 m. + # Rotor speeds differ by about 1.4% between Euler and RK4 after two steps. + np.testing.assert_allclose( + value, default_value, rtol=2e-2, atol=1e-4, err_msg=name + ) + else: + np.testing.assert_array_equal(value, default_value, err_msg=name) + else: + assert value == default_value, f"{name}: Value mismatch" + finally: + sim.close() + default_sim.close() diff --git a/tests/unit/test_sim.py b/tests/unit/test_sim.py index c3b3338d..5fe1dc9e 100644 --- a/tests/unit/test_sim.py +++ b/tests/unit/test_sim.py @@ -229,24 +229,6 @@ def test_sim_step(n_worlds: int, n_drones: int, dynamics: Dynamics, control: Con sim.step(2) -@pytest.mark.unit -@pytest.mark.parametrize("integrator", Integrator) -def test_sim_step_integrators(integrator: Integrator): - """Every integration scheme advances a default simulation without invalid state.""" - sim = Sim(integrator=integrator, device="cpu") - sim.step(2) - - assert jnp.all(sim.data.core.steps == 2) - assert jnp.all(jnp.isfinite(sim.data.states.pos)) - assert jnp.all(jnp.isfinite(sim.data.states.quat)) - assert jnp.all(jnp.isfinite(sim.data.states.vel)) - assert jnp.all(jnp.isfinite(sim.data.states.ang_vel)) - assert jnp.all(jnp.isfinite(sim.data.states.force)) - assert jnp.all(jnp.isfinite(sim.data.states.torque)) - assert jnp.all(jnp.isfinite(sim.data.states.rotor_vel)) - sim.close() - - @pytest.mark.unit def test_state_control_forwards_body_rates(): """State control must forward the body rates of the command to the attitude controller.""" From 4998870f44ac0d2393c5a8f5e4507634b56c570f Mon Sep 17 00:00:00 2001 From: Martin Schuck <57562633+amacati@users.noreply.github.com> Date: Thu, 8 Oct 2026 12:34:13 +0300 Subject: [PATCH 4/4] Simplify integrator tests --- tests/unit/test_integrators.py | 57 +++++++++++++--------------------- 1 file changed, 22 insertions(+), 35 deletions(-) diff --git a/tests/unit/test_integrators.py b/tests/unit/test_integrators.py index 5851f565..81a84266 100644 --- a/tests/unit/test_integrators.py +++ b/tests/unit/test_integrators.py @@ -15,40 +15,27 @@ def test_sim_step_integrators(integrator: Integrator, device: str): """Every integrator stays close to the default after two simulation steps.""" sim = Sim(integrator=integrator, device=device) default_sim = Sim(device=device) - try: - sim.step(2) - default_sim.step(2) + sim.step(2) + default_sim.step(2) - assert jnp.all(sim.data.core.steps == 2) - assert jnp.all(default_sim.data.core.steps == 2) - - data, structure = jax.tree.flatten_with_path(sim.data) - default_data, default_structure = jax.tree.flatten(default_sim.data) - assert structure == default_structure, "Simulation data structures must match" - for (path, value), default_value in zip(data, default_data, strict=True): - name = jax.tree_util.keystr(path) - if isinstance(value, jnp.ndarray): - assert type(value) is type(default_value), f"{name}: Type mismatch" - assert value.shape == default_value.shape, f"{name}: Shape mismatch" - assert value.dtype == default_value.dtype, f"{name}: Dtype mismatch" - assert value.device == default_value.device, f"{name}: Device mismatch" - if jax.dtypes.issubdtype(value.dtype, jax.dtypes.prng_key): - np.testing.assert_array_equal( - jax.random.key_data(value), jax.random.key_data(default_value), err_msg=name - ) - else: - assert jnp.all(jnp.isfinite(value)), f"{name}: Non-finite integrator data" - assert jnp.all(jnp.isfinite(default_value)), f"{name}: Non-finite default data" - if jnp.issubdtype(value.dtype, jnp.inexact): - # At 500 Hz, position differences are O(g * dt**2), below 1e-4 m. - # Rotor speeds differ by about 1.4% between Euler and RK4 after two steps. - np.testing.assert_allclose( - value, default_value, rtol=2e-2, atol=1e-4, err_msg=name - ) - else: - np.testing.assert_array_equal(value, default_value, err_msg=name) + data, structure = jax.tree.flatten_with_path(sim.data) + default_data, default_structure = jax.tree.flatten(default_sim.data) + assert structure == default_structure, "Simulation data structures must match" + for (path, value), default_value in zip(data, default_data, strict=True): + name = jax.tree_util.keystr(path) + if isinstance(value, jnp.ndarray): + if jax.dtypes.issubdtype(value.dtype, jax.dtypes.prng_key): + np.testing.assert_array_equal( + jax.random.key_data(value), jax.random.key_data(default_value), err_msg=name + ) + elif jnp.issubdtype(value.dtype, jnp.inexact): + # Rotor speeds differ by several percent between Euler and RK4. + np.testing.assert_allclose( + value, default_value, rtol=2e-2, atol=1e-4, err_msg=name, strict=True + ) else: - assert value == default_value, f"{name}: Value mismatch" - finally: - sim.close() - default_sim.close() + np.testing.assert_array_equal(value, default_value, err_msg=name, strict=True) + else: + assert value == default_value, f"{name}: Value mismatch" + sim.close() + default_sim.close()