Yield ownership of vector<double>, vector<float>, and vector<int> to mjSpec in Python bindings.
Fixes #2756 PiperOrigin-RevId: 785417251 Change-Id: Ib399c46ee59585258ace7d9583c68b7d06d0a9e4
This commit is contained in:
committed by
Copybara-Service
parent
60f9b34a77
commit
fc13995dd4
@@ -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(
|
||||
|
||||
@@ -221,6 +221,9 @@ PYBIND11_MODULE(_specs, m) {
|
||||
DefineArray<char>(m, "MjCharVec");
|
||||
DefineArray<std::string>(m, "MjStringVec");
|
||||
DefineArray<std::byte>(m, "MjByteVec");
|
||||
DefineArray<double>(m, "MjDoubleVec");
|
||||
DefineArray<float>(m, "MjFloatVec");
|
||||
DefineArray<int>(m, "MjIntVec");
|
||||
|
||||
// ============================= MJSPEC =====================================
|
||||
mjSpec.def(py::init<>());
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user