From da04688071a835d45ff8ff7570c97b22d8f09495 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Mon, 17 Mar 2025 09:20:29 -0700 Subject: [PATCH] Automatically compile mjSpec in to_xml(). Also, move mj_compile from MjModelWrapper to MjSpec. PiperOrigin-RevId: 737640210 Change-Id: I93d2a31c33eb359452b0d966a29939183b4c5198 --- doc/APIreference/functions.rst | 2 +- doc/APIreference/functions_override.rst | 2 +- python/mujoco/specs.cc | 72 ++++++++++++++----------- python/mujoco/specs_test.py | 8 +-- python/mujoco/structs.cc | 22 ++------ python/mujoco/structs.h | 2 +- src/xml/xml_api.cc | 7 ++- 7 files changed, 54 insertions(+), 61 deletions(-) diff --git a/doc/APIreference/functions.rst b/doc/APIreference/functions.rst index 51762bc1..6b9329a7 100644 --- a/doc/APIreference/functions.rst +++ b/doc/APIreference/functions.rst @@ -99,7 +99,7 @@ Free last XML model if loaded. Called internally at each load. .. mujoco-include:: mj_saveXMLString Save spec to XML string, return 0 on success, -1 on failure. If the length of the output buffer is too small, returns -the required size. XML saving requires that the spec first be compiled. +the required size. XML saving automatically compiles the spec before saving. .. _mj_saveXML: diff --git a/doc/APIreference/functions_override.rst b/doc/APIreference/functions_override.rst index 7278eee3..afd18be4 100644 --- a/doc/APIreference/functions_override.rst +++ b/doc/APIreference/functions_override.rst @@ -59,7 +59,7 @@ saving mechanisms. .. _mj_saveXMLString: Save spec to XML string, return 0 on success, -1 on failure. If the length of the output buffer is too small, returns -the required size. XML saving requires that the spec first be compiled. +the required size. XML saving automatically compiles the spec before saving. .. _mj_saveXML: diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 8853430e..68a5eee1 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -15,6 +15,7 @@ #include #include // IWYU pragma: keep #include +#include #include #include #include @@ -122,6 +123,41 @@ struct MjSpec { ~MjSpec() { mj_deleteSpec(ptr); } + + raw::MjModel* Compile() { + if (assets.empty()) { + auto m = mj_compile(ptr, 0); + if (!m || mjs_isWarning(ptr)) { + throw py::value_error(mjs_getError(ptr)); + } + return m; + } + mjVFS vfs; + mj_defaultVFS(&vfs); + for (const auto& asset : assets) { + std::string buffer_name = + _impl::StripPath(py::cast(asset.first).c_str()); + std::string buffer = py::cast(asset.second); + const int vfs_error = InterceptMjErrors(mj_addBufferVFS)( + &vfs, buffer_name.c_str(), buffer.c_str(), buffer.size()); + if (vfs_error) { + mj_deleteVFS(&vfs); + if (vfs_error == 2) { + throw py::value_error("Repeated file name in assets dict: " + + buffer_name); + } else { + throw py::value_error("Asset failed to load: " + buffer_name); + } + } + } + auto m = mj_compile(ptr, &vfs); + if (!m || mjs_isWarning(ptr)) { + throw py::value_error(mjs_getError(ptr)); + } + mj_deleteVFS(&vfs); + return m; + } + raw::MjSpec* ptr; py::dict assets; bool override_assets = true; @@ -245,8 +281,8 @@ void SetFrame(raw::MjsBody* body, mjtObj objtype, raw::MjsFrame* frame) { PYBIND11_MODULE(_specs, m) { auto structs_m = py::module::import("mujoco._structs"); - py::function mjmodel_from_spec_ptr = - structs_m.attr("MjModel").attr("_from_spec_ptr"); + py::function mjmodel_from_raw_ptr = + structs_m.attr("MjModel").attr("_from_model_ptr"); py::function mjmodel_mjdata_from_spec_ptr = structs_m.attr("_recompile_spec_addr"); @@ -416,33 +452,8 @@ PYBIND11_MODULE(_specs, m) { return mjs_findDefault(self.ptr, classname.c_str()); }, py::return_value_policy::reference_internal); - mjSpec.def("compile", [mjmodel_from_spec_ptr](MjSpec& self) -> py::object { - if (self.assets.empty()) { - return mjmodel_from_spec_ptr(reinterpret_cast(self.ptr)); - } - mjVFS vfs; - mj_defaultVFS(&vfs); - for (const auto& asset : self.assets) { - std::string buffer_name = - _impl::StripPath(py::cast(asset.first).c_str()); - std::string buffer = py::cast(asset.second); - const int vfs_error = InterceptMjErrors(mj_addBufferVFS)( - &vfs, buffer_name.c_str(), buffer.c_str(), buffer.size()); - if (vfs_error) { - mj_deleteVFS(&vfs); - if (vfs_error == 2) { - throw py::value_error("Repeated file name in assets dict: " + - buffer_name); - } else { - throw py::value_error("Asset failed to load: " + buffer_name); - } - } - } - auto model = - mjmodel_from_spec_ptr(reinterpret_cast(self.ptr), - reinterpret_cast(&vfs)); - mj_deleteVFS(&vfs); - return model; + mjSpec.def("compile", [mjmodel_from_raw_ptr](MjSpec& self) -> py::object { + return mjmodel_from_raw_ptr(reinterpret_cast(self.Compile())); }); mjSpec.def_property( "assets", @@ -463,9 +474,10 @@ PYBIND11_MODULE(_specs, m) { self.override_assets = override_assets; }); mjSpec.def("to_xml", [](MjSpec& self) -> std::string { + mj_deleteModel(self.Compile()); + std::array err; int size = mj_saveXMLString(self.ptr, nullptr, 0, nullptr, 0); std::unique_ptr buf(new char[size + 1]); - std::array err; buf[0] = '\0'; err[0] = '\0'; mj_saveXMLString(self.ptr, buf.get(), size + 1, err.data(), err.size()); diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 83493481..46bb2af9 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -520,13 +520,9 @@ class SpecsTest(absltest.TestCase): np.testing.assert_array_equal(data_new.qpos[:4], data.qpos) np.testing.assert_array_equal(data_new.qvel[:3], data.qvel) - def test_uncompiled_spec_cannot_be_written(self): + def test_uncompiled_spec_can_be_written(self): spec = mujoco.MjSpec() - - # Cannot write XML of an uncompiled spec. - expected_error = 'XML Write error: Only compiled model can be written' - with self.assertRaisesWithLiteralMatch(mujoco.FatalError, expected_error): - spec.to_xml() + spec.to_xml() def test_modelname_default_class(self): XML = textwrap.dedent("""\ diff --git a/python/mujoco/structs.cc b/python/mujoco/structs.cc index 6112653e..f9272721 100644 --- a/python/mujoco/structs.cc +++ b/python/mujoco/structs.cc @@ -408,12 +408,7 @@ MjModelWrapper MjModelWrapper::LoadXML( return MjModelWrapper(model); } -MjModelWrapper MjModelWrapper::CompileSpec(raw::MjSpec* spec, - const mjVFS* vfs) { - auto m = mj_compile(spec, vfs); - if (!m || mjs_isWarning(spec)) { - throw py::value_error(mjs_getError(spec)); - } +MjModelWrapper MjModelWrapper::WrapRawModel(raw::MjModel* m) { return MjModelWrapper(m); } @@ -1616,18 +1611,9 @@ PYBIND11_MODULE(_structs, m) { py::arg("xml"), py::arg_v("assets", py::none()), py::doc( R"(Loads an MjModel from an XML string and an optional assets dictionary.)")); - mjModel.def_static( - "_from_spec_ptr", [](uintptr_t addr) { - return MjModelWrapper::CompileSpec( - reinterpret_cast(addr), - nullptr); - }); - mjModel.def_static( - "_from_spec_ptr", [](uintptr_t addr, uintptr_t vfs) { - return MjModelWrapper::CompileSpec( - reinterpret_cast(addr), - reinterpret_cast(vfs)); - }); + mjModel.def_static("_from_model_ptr", [](uintptr_t addr) { + return MjModelWrapper::WrapRawModel(reinterpret_cast(addr)); + }); mjModel.def_static( "from_xml_path", &MjModelWrapper::LoadXMLFile, py::arg("filename"), py::arg_v("assets", py::none()), diff --git a/python/mujoco/structs.h b/python/mujoco/structs.h index 71e52190..418bb879 100644 --- a/python/mujoco/structs.h +++ b/python/mujoco/structs.h @@ -533,7 +533,7 @@ class MjWrapper : public WrapperBase { const std::optional< std::unordered_map>& assets); - static MjWrapper CompileSpec(raw::MjSpec* spec, const mjVFS* vfs); + static MjWrapper WrapRawModel(raw::MjModel* m); static constexpr char kFromRawPointer[] = "__MUJOCO_STRUCTS_MJMODELWRAPPER_LOOKUP"; diff --git a/src/xml/xml_api.cc b/src/xml/xml_api.cc index 627e2699..6952836f 100644 --- a/src/xml/xml_api.cc +++ b/src/xml/xml_api.cc @@ -246,15 +246,14 @@ int mj_saveXML(const mjSpec* s, const char* filename, char* error, int error_sz) // if length of the output buffer is too small, returns the required size int mj_saveXMLString(const mjSpec* s, char* xml, int xml_sz, char* error, int error_sz) { std::string result = WriteXML(NULL, s, error, error_sz); - if (result.size() >= xml_sz) { + if (result.empty()) { + return -1; + } else if (result.size() >= xml_sz) { std::string error_msg = "Output string too short, should be at least " + std::to_string(result.size()+1); mjCopyError(error, error_msg.c_str(), error_sz); return result.size(); } - if (result.empty()) { - return -1; - } result.copy(xml, xml_sz); xml[result.size()] = 0;