From f49ca01bef00855fa0744ab03dc6b8eb71eccfa4 Mon Sep 17 00:00:00 2001 From: Martin Schuck Date: Wed, 16 Sep 2026 13:14:08 +0200 Subject: [PATCH 1/2] Avoid materializing init data on single device --- crazyflow/sim/sharding.py | 25 +++++++++++++++++++++++-- crazyflow/sim/sim.py | 34 ++++++++++++++++++++++------------ 2 files changed, 45 insertions(+), 14 deletions(-) diff --git a/crazyflow/sim/sharding.py b/crazyflow/sim/sharding.py index 2312a7c7..df240575 100644 --- a/crazyflow/sim/sharding.py +++ b/crazyflow/sim/sharding.py @@ -15,9 +15,9 @@ from crazyflow.utils import world_mask if TYPE_CHECKING: - from collections.abc import Sequence + from collections.abc import Callable, Sequence - from jax import Device + from jax import Array, Device from jax.sharding import Mesh from crazyflow.sim.data import SimData @@ -71,3 +71,24 @@ def shard(data: SimData, mesh: Mesh) -> SimData: The placed simulation data. """ return jax.device_put(data, placement(data, mesh)) + + +def build_sharded( + create: Callable[[int | Array], SimData], rng_key: int | Array, mesh: Mesh +) -> SimData: + """Build simulation data distributed over a mesh. + + Tracing the construction lets us put the result directly on the mesh and prevents data from + being materialised on a single device. + + Args: + create: Callable that builds the data from an rng key. + rng_key: Random number generator key for the simulation, or a seed to derive one from. + mesh: Mesh to distribute the worlds over. + + Returns: + The placed simulation data. + """ + if isinstance(rng_key, int): # Tracing turns a seed into an array that is not a key + rng_key = jax.random.key(rng_key) + return jax.jit(create, out_shardings=placement(jax.eval_shape(create, rng_key), mesh))(rng_key) diff --git a/crazyflow/sim/sim.py b/crazyflow/sim/sim.py index 055f8a54..96db91e8 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -35,7 +35,7 @@ from crazyflow.sim.data import SimControls, SimCore, SimData, SimParams, SimState, SimStateDeriv from crazyflow.sim.integration import Integrator, euler, rk4, symplectic_euler from crazyflow.sim.pipeline import append_fn -from crazyflow.sim.sharding import placement +from crazyflow.sim.sharding import build_sharded, placement from crazyflow.utils import grid_2d, pytree_replace, world_mask if TYPE_CHECKING: @@ -89,6 +89,7 @@ def __init__( xml_path: Path | None = None, rng_key: int = 0, fused_mjx_model: bool = False, + mesh: Mesh | None = None, ): """Build the scene and the step and reset pipelines, and allocate the batched sim data. @@ -110,6 +111,7 @@ def __init__( fused_mjx_model: If True, use the ``drone_fused`` body whose visual geometry is fused into a single mesh. This shrinks the MJX model and reduces its memory footprint at the cost of visual detail. + mesh: Mesh to distribute the worlds over. """ assert Dynamics(dynamics) in Dynamics, f"Dynamics mode {dynamics} not implemented" assert Control(control) in Control, f"Control mode {control} not implemented" @@ -123,6 +125,7 @@ def __init__( self.drone = drone self.integrator = integrator self.device = jax.devices(device)[0] + self.mesh = mesh self.n_worlds = n_worlds self.n_drones = n_drones self.freq = freq @@ -136,9 +139,11 @@ def __init__( self.mj_model, self.mj_data, self.mjx_model, self.mjx_data = self.build_mjx_model(self.spec) self.viewer: MujocoRenderer | None = None - self.data = self.init_data( - state_freq, attitude_freq, body_rate_freq, force_torque_freq, rng_key - ) + freqs = (state_freq, attitude_freq, body_rate_freq, force_torque_freq) + if mesh is None: + self.data = self.init_data(*freqs, rng_key) + else: + self.data = build_sharded(partial(self.init_data, *freqs), rng_key, mesh) self.default_data: SimData = self.build_default_data() # Build the simulation pipeline and overwrite the default _step implementation with it @@ -444,9 +449,12 @@ def build_data(self) -> SimData: attitude_freq = 0 if (a := self.data.controls.attitude) is None else a.freq body_rate_freq = 0 if (br := self.data.controls.body_rate) is None else br.freq force_torque_freq = 0 if (ft := self.data.controls.force_torque) is None else ft.freq - self.data = self.init_data( - state_freq, attitude_freq, body_rate_freq, force_torque_freq, self.data.core.rng_key - ) + freqs = (state_freq, attitude_freq, body_rate_freq, force_torque_freq) + rng_key = self.data.core.rng_key + if self.mesh is None: + self.data = self.init_data(*freqs, rng_key) + else: + self.data = build_sharded(partial(self.init_data, *freqs), rng_key, self.mesh) return self.data def shard(self, mesh: Mesh) -> SimData: @@ -459,6 +467,7 @@ def shard(self, mesh: Mesh) -> SimData: Returns: The placed simulation data. """ + self.mesh = mesh self.data = jax.device_put(self.data, placement(self.data, mesh)) self.default_data = jax.device_put(self.default_data, placement(self.default_data, mesh)) return self.data @@ -493,14 +502,15 @@ def init_data( rng_key: Array, ) -> SimData: """Initialize the simulation data.""" + device = self.device if self.mesh is None else None # Sharded data is placed by the caller drone_name = "drone_fused" if self.fused_mjx_model else "drone" drone_mocap_ids = [ self.mj_model.body(f"{drone_name}:{i}").mocapid.item() for i in range(self.n_drones) ] N, D = self.n_worlds, self.n_drones data = SimData( - states=SimState.create(N, D, self.device), - states_deriv=SimStateDeriv.create(N, D, self.device), + states=SimState.create(N, D, device), + states_deriv=SimStateDeriv.create(N, D, device), controls=SimControls.create( N, D, @@ -510,10 +520,10 @@ def init_data( attitude_freq, body_rate_freq, force_torque_freq, - self.device, + device, ), - params=SimParams.create(self.dynamics, self.drone, self.device), - core=SimCore.create(self.freq, N, D, drone_mocap_ids, rng_key, self.device), + params=SimParams.create(self.dynamics, self.drone, device), + core=SimCore.create(self.freq, N, D, drone_mocap_ids, rng_key, device), ) if D > 1: # If multiple drones, arrange them in a grid grid = grid_2d(D) From 7ebbbf6d3cc542ba269fb2a96702b4317bf7657d Mon Sep 17 00:00:00 2001 From: Martin Schuck <57562633+amacati@users.noreply.github.com> Date: Thu, 17 Sep 2026 00:27:46 +0200 Subject: [PATCH 2/2] Apply suggestions from code review Co-authored-by: Marcel Rath <75042654+ratheron@users.noreply.github.com> --- crazyflow/sim/sharding.py | 2 +- crazyflow/sim/sim.py | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/crazyflow/sim/sharding.py b/crazyflow/sim/sharding.py index df240575..f6d32200 100644 --- a/crazyflow/sim/sharding.py +++ b/crazyflow/sim/sharding.py @@ -73,7 +73,7 @@ def shard(data: SimData, mesh: Mesh) -> SimData: return jax.device_put(data, placement(data, mesh)) -def build_sharded( +def build_sharded_data( create: Callable[[int | Array], SimData], rng_key: int | Array, mesh: Mesh ) -> SimData: """Build simulation data distributed over a mesh. diff --git a/crazyflow/sim/sim.py b/crazyflow/sim/sim.py index 96db91e8..77cad553 100644 --- a/crazyflow/sim/sim.py +++ b/crazyflow/sim/sim.py @@ -35,7 +35,7 @@ from crazyflow.sim.data import SimControls, SimCore, SimData, SimParams, SimState, SimStateDeriv from crazyflow.sim.integration import Integrator, euler, rk4, symplectic_euler from crazyflow.sim.pipeline import append_fn -from crazyflow.sim.sharding import build_sharded, placement +from crazyflow.sim.sharding import build_sharded_data, placement from crazyflow.utils import grid_2d, pytree_replace, world_mask if TYPE_CHECKING: @@ -143,7 +143,7 @@ def __init__( if mesh is None: self.data = self.init_data(*freqs, rng_key) else: - self.data = build_sharded(partial(self.init_data, *freqs), rng_key, mesh) + self.data = build_sharded_data(partial(self.init_data, *freqs), rng_key, mesh) self.default_data: SimData = self.build_default_data() # Build the simulation pipeline and overwrite the default _step implementation with it @@ -454,7 +454,7 @@ def build_data(self) -> SimData: if self.mesh is None: self.data = self.init_data(*freqs, rng_key) else: - self.data = build_sharded(partial(self.init_data, *freqs), rng_key, self.mesh) + self.data = build_sharded_data(partial(self.init_data, *freqs), rng_key, self.mesh) return self.data def shard(self, mesh: Mesh) -> SimData: