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:
committed by
Copybara-Service
parent
35774706ee
commit
da04688071
+42
-30
@@ -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());
|
||||
|
||||
@@ -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("""\
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -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";
|
||||
|
||||
Reference in New Issue
Block a user