diff --git a/python/mujoco/codegen/generate_spec_bindings.py b/python/mujoco/codegen/generate_spec_bindings.py index e7b13691..6507b764 100644 --- a/python/mujoco/codegen/generate_spec_bindings.py +++ b/python/mujoco/codegen/generate_spec_bindings.py @@ -193,7 +193,7 @@ def _ptr_binding_code( []({rawclassname}& self, std::string_view {varname}) {{ *(self.{fullvarname}) = {varname}; }});""" - elif ( # C++ vectors of values -> Python array + elif ( # C++ vectors of values -> custom array vartype == 'mjDoubleVec' or vartype == 'mjFloatVec' or vartype == 'mjIntVec' @@ -202,9 +202,9 @@ def _ptr_binding_code( return f"""\ {classname}.def_property( "{varname}", - []({rawclassname}& self) -> py::array_t<{vartype}> {{ - return py::array_t<{vartype}>(self.{fullvarname}->size(), - self.{fullvarname}->data()); + []({rawclassname}& self) -> MjTypeVec<{vartype}> {{ + return MjTypeVec<{vartype}>(self.{fullvarname}->data(), + self.{fullvarname}->size()); }}, []({rawclassname}& self, py::object rhs) {{ self.{fullvarname}->clear(); @@ -212,7 +212,7 @@ def _ptr_binding_code( for (auto val : rhs) {{ self.{fullvarname}->push_back(py::cast<{vartype}>(val)); }} - }}, py::return_value_policy::reference_internal);""" + }}, py::return_value_policy::move);""" elif vartype == 'mjByteVec': return f"""\ {classname}.def_property( diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index a1ed3780..4fec07e2 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -221,6 +221,9 @@ PYBIND11_MODULE(_specs, m) { DefineArray(m, "MjCharVec"); DefineArray(m, "MjStringVec"); DefineArray(m, "MjByteVec"); + DefineArray(m, "MjDoubleVec"); + DefineArray(m, "MjFloatVec"); + DefineArray(m, "MjIntVec"); // ============================= MJSPEC ===================================== mjSpec.def(py::init<>()); diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 61b3b745..8050a745 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -98,9 +98,13 @@ class SpecsTest(absltest.TestCase): site = body.add_site() site.name = 'sitename' site.type = mujoco.mjtGeom.mjGEOM_BOX - site.userdata = [1, 2, 3, 4, 5, 6] + site.userdata = [7, 2, 3, 4, 5, 6] self.assertEqual(site.name, 'sitename') self.assertEqual(site.type, mujoco.mjtGeom.mjGEOM_BOX) + np.testing.assert_array_equal(site.userdata, [7, 2, 3, 4, 5, 6]) + + # Modify a single element of userdata. + site.userdata[0] = 1 np.testing.assert_array_equal(site.userdata, [1, 2, 3, 4, 5, 6]) # Compile the spec and check for expected values in the model.