_realloc_island

PiperOrigin-RevId: 941618362
Change-Id: Idef1deb0ff4976de42aa585993988ae3363cd1fc
This commit is contained in:
Taylor Howell
2026-07-02 04:19:19 -07:00
committed by Copybara-Service
parent 0e58c48504
commit c3adb50898
2 changed files with 112 additions and 0 deletions
+43
View File
@@ -699,6 +699,49 @@ class MuJoCoBindingsTest(parameterized.TestCase):
self.assertEmpty(self.data.contact)
self.assertEmpty(self.data.efc_id)
def test_realloc_island(self):
# Test allocation on fresh data (on its own)
nisland = 2
nidof = 4
mujoco._functions._realloc_island(self.data, nisland=nisland, nidof=nidof)
self.assertEqual(self.data.nisland, nisland)
self.assertEqual(self.data.nidof, nidof)
self.assertEqual(self.data.island_nv.shape, (nisland,))
self.assertEqual(self.data.ifrc_smooth.shape, (nidof,))
# Test allocation after _realloc_con_efc
nefc = 10
mujoco._functions._realloc_con_efc(self.data, ncon=0, nefc=nefc)
nisland = 3
nidof = 5
mujoco._functions._realloc_island(self.data, nisland=nisland, nidof=nidof)
self.assertEqual(self.data.nisland, nisland)
self.assertEqual(self.data.nidof, nidof)
self.assertEqual(self.data.island_nv.shape, (nisland,))
self.assertEqual(self.data.ifrc_smooth.shape, (nidof,))
# Test re-allocation (calling it again with different sizes)
nisland2 = 4
nidof2 = 6
mujoco._functions._realloc_island(self.data, nisland=nisland2, nidof=nidof2)
self.assertEqual(self.data.nisland, nisland2)
self.assertEqual(self.data.nidof, nidof2)
self.assertEqual(self.data.island_nv.shape, (nisland2,))
self.assertEqual(self.data.ifrc_smooth.shape, (nidof2,))
# Test insufficient memory handling
expected_error = (
r'Insufficient arena memory, currently allocated memory=' +
r'"[0-9]+[A-Z]?". Increase using <size memory="X"/>.'
)
with self.assertRaisesRegex(mujoco.FatalError, expected_error):
mujoco._functions._realloc_island(self.data, 100000000, 100000000)
self.assertEqual(self.data.nisland, 0)
self.assertEqual(self.data.nidof, 0)
def test_mj_struct_list_equality(self):
model2 = mujoco.MjModel.from_xml_string(TEST_XML)
data2 = mujoco.MjData(model2)
+69
View File
@@ -1696,6 +1696,7 @@ PYBIND11_MODULE(_functions, pymodule, pybind11::mod_gil_not_used()) {
#define X(type, name, nr, nc) data->name = nullptr;
MJDATA_ARENA_POINTERS_SOLVER
MJDATA_ARENA_POINTERS_DUAL
MJDATA_ARENA_POINTERS_ISLAND
#undef X
};
@@ -1744,6 +1745,74 @@ PYBIND11_MODULE(_functions, pymodule, pybind11::mod_gil_not_used()) {
},
py::arg("d"), py::arg("ncon"), py::arg("nefc"), py::arg("nJ") = -1,
py::call_guard<py::gil_scoped_release>());
pymodule.def(
"_realloc_island",
[](MjDataWrapper& d, int nisland, int nidof) {
raw::MjData* data = d.get();
size_t parena_start = data->parena;
// Find island block start in arena to reclaim memory on re-allocation.
char* min_ptr = nullptr;
#define X(type, name, nr, nc) \
if (data->name && \
(!min_ptr || reinterpret_cast<char*>(data->name) < min_ptr)) { \
min_ptr = reinterpret_cast<char*>(data->name); \
}
MJDATA_ARENA_POINTERS_ISLAND
#undef X
if (min_ptr && data->arena) {
parena_start = min_ptr - static_cast<char*>(data->arena);
}
auto cleanup = [](raw::MjData* data, size_t target_parena) {
#define X(type, name, nr, nc) data->name = nullptr;
MJDATA_ARENA_POINTERS_ISLAND
#undef X
data->nisland = 0;
data->nidof = 0;
data->parena = target_parena;
#ifdef ADDRESS_SANITIZER
ASAN_POISON_MEMORY_REGION(
static_cast<char*>(data->arena) + target_parena,
data->narena - data->pstack - target_parena);
#endif
};
cleanup(data, parena_start);
char error_msg[128];
error_msg[0] = '\0';
const char* error_msg_fmt =
"Insufficient arena memory, currently allocated memory=\"%s\". "
"Increase using <size memory=\"X\"/>.";
data->nisland = nisland;
data->nidof = nidof;
#undef MJ_M
#define MJ_M(x) d.model().get()->x
#undef MJ_D
#define MJ_D(x) data->x
#define X(type, name, nr, nc) \
data->name = static_cast<type*>(InterceptMjErrors(::mj_arenaAllocByte)( \
data, sizeof(type) * (nr) * (nc), alignof(type))); \
if (!data->name) { \
cleanup(data, parena_start); \
std::snprintf(error_msg, sizeof(error_msg), error_msg_fmt, \
mju_writeNumBytes(data->narena)); \
throw FatalError(error_msg); \
}
MJDATA_ARENA_POINTERS_ISLAND
#undef X
#undef MJ_D
#define MJ_D(x) x
#undef MJ_M
#define MJ_M(x) x
},
py::arg("d"), py::arg("nisland"), py::arg("nidof"),
py::call_guard<py::gil_scoped_release>());
} // PYBIND11_MODULE NOLINT(readability/fn_size)
} // namespace
} // namespace mujoco::python