diff --git a/mujoco_warp/_src/io.py b/mujoco_warp/_src/io.py index a8c00483f..fa98a6f54 100644 --- a/mujoco_warp/_src/io.py +++ b/mujoco_warp/_src/io.py @@ -2176,6 +2176,28 @@ def get_data_into( # TODO(team): if sparse, set nJ based on sparse efc_J mujoco._functions._realloc_con_efc(result, ncon=ncon, nefc=nefc, nJ=nefc * mjm.nv) + # Check island compatibility between MjModel and Data + mjm_islands_enabled = not bool(mjm.opt.disableflags & mujoco.mjtDisableBit.mjDSBL_ISLAND) + + # Check if device island arrays are allocated (automatically enabled when sleep is enabled). + # TODO(team): update when islands are enabled by default. + d_islands_enabled = d.tree_island.shape[1] > 0 + + if mjm_islands_enabled != d_islands_enabled: + raise ValueError("MjModel and Data island configurations do not match.") + + # Reallocate host island fields if islands/sleep are enabled in host model options. + host_islands_allocated = mjm_islands_enabled and ( + bool(mjm.opt.enableflags & mujoco.mjtEnableBit.mjENBL_SLEEP) + or bool(mjm.opt.enableflags & getattr(mujoco.mjtEnableBit, "mjENBL_ISLAND", 0)) + ) + + if host_islands_allocated: + nisland = d.nisland.numpy()[world_id] + nidof = d.nidof.numpy()[world_id] + if result.island_idofadr.shape[0] < nisland or result.dof_island.shape[0] < nidof: + mujoco._functions._realloc_island(result, nisland=nisland, nidof=nidof) + ne = d.ne.numpy()[world_id] nf = d.nf.numpy()[world_id] nl = d.nl.numpy()[world_id] @@ -2370,23 +2392,29 @@ def get_data_into( result.body_awake[:] = d.body_awake.numpy()[world_id] # islands - nisland = d.nisland.numpy()[world_id] - result.nisland = nisland - if d.tree_island.shape[1] > 0 and nisland: - result.tree_island[:] = d.tree_island.numpy()[world_id] - result.dof_island[:] = d.dof_island.numpy()[world_id] - result.island_idofadr[:nisland] = d.island_idofadr.numpy()[world_id, :nisland] - result.island_dofadr[:nisland] = d.island_dofadr.numpy()[world_id, :nisland] - result.island_nv[:nisland] = d.island_nv.numpy()[world_id, :nisland] - result.island_nefc[:nisland] = d.island_nefc.numpy()[world_id, :nisland] - result.island_ne[:nisland] = d.island_ne.numpy()[world_id, :nisland] - result.island_nf[:nisland] = d.island_nf.numpy()[world_id, :nisland] - result.island_iefcadr[:nisland] = d.island_iefcadr.numpy()[world_id, :nisland] - nv = mjm.nv - result.map_dof2idof[:nv] = d.map_dof2idof.numpy()[world_id, :nv] - result.map_idof2dof[:nv] = d.map_idof2dof.numpy()[world_id, :nv] - result.map_efc2iefc[:nefc] = d.map_efc2iefc.numpy()[world_id, :nefc] - result.map_iefc2efc[:nefc] = d.map_iefc2efc.numpy()[world_id, :nefc] + if mjm_islands_enabled and d.tree_island.shape[1] > 0: + nisland = d.nisland.numpy()[world_id] + result.nisland = nisland + if nisland > 0: + result.tree_island[:] = d.tree_island.numpy()[world_id] + result.dof_island[:] = d.dof_island.numpy()[world_id] + if result.island_idofadr.shape[0] >= nisland: + result.island_idofadr[:nisland] = d.island_idofadr.numpy()[world_id, :nisland] + result.island_dofadr[:nisland] = d.island_dofadr.numpy()[world_id, :nisland] + result.island_nv[:nisland] = d.island_nv.numpy()[world_id, :nisland] + result.island_nefc[:nisland] = d.island_nefc.numpy()[world_id, :nisland] + result.island_ne[:nisland] = d.island_ne.numpy()[world_id, :nisland] + result.island_nf[:nisland] = d.island_nf.numpy()[world_id, :nisland] + result.island_iefcadr[:nisland] = d.island_iefcadr.numpy()[world_id, :nisland] + nv = mjm.nv + if result.map_dof2idof.shape[0] >= nv: + result.map_dof2idof[:nv] = d.map_dof2idof.numpy()[world_id, :nv] + result.map_idof2dof[:nv] = d.map_idof2dof.numpy()[world_id, :nv] + if result.map_efc2iefc.shape[0] >= nefc: + result.map_efc2iefc[:nefc] = d.map_efc2iefc.numpy()[world_id, :nefc] + result.map_iefc2efc[:nefc] = d.map_iefc2efc.numpy()[world_id, :nefc] + else: + result.nisland = 0 def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): diff --git a/mujoco_warp/_src/io_test.py b/mujoco_warp/_src/io_test.py index 58ed68145..7cfdbfa3f 100644 --- a/mujoco_warp/_src/io_test.py +++ b/mujoco_warp/_src/io_test.py @@ -747,6 +747,63 @@ def test_get_data_into(self, nworld, world_id): field, ) + def test_get_data_into_realloc_island(self): + mjm, mjd, m, d = test_data.fixture( + xml=""" + + + + + + + + + + + + + + + """, + nworld=1, + ) + + expected_nisland = int(d.nisland.numpy()[0]) + expected_nidof = int(d.nidof.numpy()[0]) + + mjwarp.get_data_into(mjd, mjm, d, world_id=0) + + self.assertEqual(mjd.nisland, expected_nisland) + self.assertEqual(mjd.island_idofadr.shape[0], expected_nisland) + + def test_get_data_into_islands_disabled_error(self): + mjm, mjd, m, d = test_data.fixture( + xml=""" + + + + + + + + + + + + + + """, + nworld=1, + ) + mjwarp.forward(m, d) + if d.nisland.numpy()[0] > 0: + with self.assertRaisesRegex(ValueError, "islands are disabled in MjModel"): + mjwarp.get_data_into(mjd, mjm, d, world_id=0) + @parameterized.product( xml=_IO_TEST_MODELS, cone=list(ConeType),