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
25 changes: 23 additions & 2 deletions crazyflow/sim/sharding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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_data(
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)
34 changes: 22 additions & 12 deletions crazyflow/sim/sim.py
Original file line number Diff line number Diff line change
Expand Up @@ -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_data, placement
from crazyflow.utils import grid_2d, pytree_replace, world_mask

if TYPE_CHECKING:
Expand Down Expand Up @@ -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.

Expand All @@ -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"
Expand All @@ -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
Expand All @@ -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_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
Expand Down Expand Up @@ -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_data(partial(self.init_data, *freqs), rng_key, self.mesh)
return self.data

def shard(self, mesh: Mesh) -> SimData:
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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)
Expand Down
Loading