Automatically compile mjSpec in to_xml().

Also, move mj_compile from MjModelWrapper to MjSpec.

PiperOrigin-RevId: 737640210
Change-Id: I93d2a31c33eb359452b0d966a29939183b4c5198
This commit is contained in:
Alessio Quaglino
2025-03-17 09:20:29 -07:00
committed by Copybara-Service
parent 35774706ee
commit da04688071
7 changed files with 54 additions and 61 deletions
+1 -1
View File
@@ -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:
+1 -1
View File
@@ -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:
+42 -30
View File
@@ -15,6 +15,7 @@
#include <array>
#include <cstddef> // IWYU pragma: keep
#include <cstdint>
#include <cstring>
#include <memory>
#include <optional>
#include <string>
@@ -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<std::string>(asset.first).c_str());
std::string buffer = py::cast<std::string>(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<uintptr_t>(self.ptr));
}
mjVFS vfs;
mj_defaultVFS(&vfs);
for (const auto& asset : self.assets) {
std::string buffer_name =
_impl::StripPath(py::cast<std::string>(asset.first).c_str());
std::string buffer = py::cast<std::string>(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<uintptr_t>(self.ptr),
reinterpret_cast<uintptr_t>(&vfs));
mj_deleteVFS(&vfs);
return model;
mjSpec.def("compile", [mjmodel_from_raw_ptr](MjSpec& self) -> py::object {
return mjmodel_from_raw_ptr(reinterpret_cast<uintptr_t>(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<char, 1024> err;
int size = mj_saveXMLString(self.ptr, nullptr, 0, nullptr, 0);
std::unique_ptr<char[]> buf(new char[size + 1]);
std::array<char, 1024> err;
buf[0] = '\0';
err[0] = '\0';
mj_saveXMLString(self.ptr, buf.get(), size + 1, err.data(), err.size());
+2 -6
View File
@@ -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("""\
+4 -18
View File
@@ -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<raw::MjSpec*>(addr),
nullptr);
});
mjModel.def_static(
"_from_spec_ptr", [](uintptr_t addr, uintptr_t vfs) {
return MjModelWrapper::CompileSpec(
reinterpret_cast<raw::MjSpec*>(addr),
reinterpret_cast<mjVFS*>(vfs));
});
mjModel.def_static("_from_model_ptr", [](uintptr_t addr) {
return MjModelWrapper::WrapRawModel(reinterpret_cast<raw::MjModel*>(addr));
});
mjModel.def_static(
"from_xml_path", &MjModelWrapper::LoadXMLFile,
py::arg("filename"), py::arg_v("assets", py::none()),
+1 -1
View File
@@ -533,7 +533,7 @@ class MjWrapper<raw::MjModel> : public WrapperBase<raw::MjModel> {
const std::optional<
std::unordered_map<std::string, pybind11::bytes>>& assets);
static MjWrapper CompileSpec(raw::MjSpec* spec, const mjVFS* vfs);
static MjWrapper WrapRawModel(raw::MjModel* m);
static constexpr char kFromRawPointer[] =
"__MUJOCO_STRUCTS_MJMODELWRAPPER_LOOKUP";
+3 -4
View File
@@ -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;