diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 4b69b9e4..4989ee26 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -8017,6 +8017,15 @@ void mj_saveModel_wrapper(const MjModel& m, const StringOrNull& filename, const mj_saveModel(m.get(), filename_.data(), buffer_.data(), static_cast(buffer_.size())); } +std::unique_ptr mj_loadModel_wrapper(std::string filename, const MjVFS& vfs) { + mjModel *model = mj_loadModel(filename.c_str(), vfs.get()); + if (!model) { + printf("mj_loadModel: failed to load from mjb"); + return nullptr; + } + return std::unique_ptr(new MjModel(model)); +} + std::unique_ptr parseXMLString_wrapper(const std::string &xml) { char error[1000]; mjSpec *ptr = mj_parseXMLString(xml.c_str(), nullptr, error, sizeof(error)); @@ -11206,6 +11215,7 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) { .property("uselimit", &MjLROpt::uselimit, &MjLROpt::set_uselimit, reference()); emscripten::class_("MjModel") .class_function("mj_loadXML", &mj_loadXML_wrapper, take_ownership()) + .class_function("mj_loadBinary", &mj_loadModel_wrapper, take_ownership()) .constructor() // Binds the functions on MjModel that return accessors. #define X_ACCESSOR(NAME, Name, OBJTYPE, field_name, nfield) \ diff --git a/wasm/codegen/generators/constants.py b/wasm/codegen/generators/constants.py index 5046ab1e..dedbf5a6 100644 --- a/wasm/codegen/generators/constants.py +++ b/wasm/codegen/generators/constants.py @@ -74,6 +74,7 @@ _SKIPPED_CLASS_METHODS: tuple[str, ...] = ( "mj_deleteModel", "mj_deleteSpec", "mj_deleteVFS", + "mj_loadModel", "mj_loadXML", "mj_makeData", "mj_makeSpec", @@ -139,7 +140,6 @@ _SKIPPED_MEMORY_FUNCTIONS: tuple[str, ...] = ( # go/keep-sorted start "mj_freeLastXML", "mj_freeStack", - "mj_loadModel", "mj_loadModelBuffer", "mj_markStack", "mj_stackAllocByte", diff --git a/wasm/codegen/generators/structs.py b/wasm/codegen/generators/structs.py index 9eafd731..ce4659b3 100644 --- a/wasm/codegen/generators/structs.py +++ b/wasm/codegen/generators/structs.py @@ -508,10 +508,14 @@ def _build_struct_bindings( builder.line(".constructor()") builder.line(".constructor()") elif w == "MjModel": - fn = common.wrapped_function_name( + f1 = common.wrapped_function_name( introspect_functions.FUNCTIONS["mj_loadXML"] ) - builder.line(f'.class_function("mj_loadXML", &{fn}, take_ownership())') + builder.line(f'.class_function("mj_loadXML", &{f1}, take_ownership())') + f2 = common.wrapped_function_name( + introspect_functions.FUNCTIONS["mj_loadModel"] + ) + builder.line(f'.class_function("mj_loadBinary", &{f2}, take_ownership())') builder.line(".constructor()") builder.line(""" // Binds the functions on MjModel that return accessors. diff --git a/wasm/codegen/templates/bindings.cc b/wasm/codegen/templates/bindings.cc index 88a640ff..14b0457c 100644 --- a/wasm/codegen/templates/bindings.cc +++ b/wasm/codegen/templates/bindings.cc @@ -535,6 +535,15 @@ void mj_saveModel_wrapper(const MjModel& m, const StringOrNull& filename, const mj_saveModel(m.get(), filename_.data(), buffer_.data(), static_cast(buffer_.size())); } +std::unique_ptr mj_loadModel_wrapper(std::string filename, const MjVFS& vfs) { + mjModel *model = mj_loadModel(filename.c_str(), vfs.get()); + if (!model) { + printf("mj_loadModel: failed to load from mjb"); + return nullptr; + } + return std::unique_ptr(new MjModel(model)); +} + std::unique_ptr parseXMLString_wrapper(const std::string &xml) { char error[1000]; mjSpec *ptr = mj_parseXMLString(xml.c_str(), nullptr, error, sizeof(error)); diff --git a/wasm/tests/bindings_test.ts b/wasm/tests/bindings_test.ts index c46d79a5..faa3dc3e 100644 --- a/wasm/tests/bindings_test.ts +++ b/wasm/tests/bindings_test.ts @@ -2176,4 +2176,88 @@ describe('MuJoCo WASM Bindings', () => { unlinkXMLFile(filename); } }); + + it('should load and save a model with assets to binary', () => { + const xmlContent = ` + + + + + + + + `; + const cube1 = ` + v -1 -1 1 + v 1 -1 1 + v -1 1 1 + v 1 1 1 + v -1 1 -1 + v 1 1 -1 + v -1 -1 -1 + v 1 -1 -1`; + const xmlFilename = '/tmp/binary_test.xml'; + const objFilename = '/tmp/cube.obj'; + const mjbFilename = '/tmp/binary_test.mjb'; + + writeXMLFile(xmlFilename, xmlContent); + writeXMLFile(objFilename, cube1); + + let model: MjModel|null = null; + let binaryModel: MjModel|null = null; + let vfs: MjVFS|null = null; + + try { + model = mujoco.MjModel.mj_loadXML(xmlFilename); + assertExists(model); + + mujoco.mj_saveModel(model, mjbFilename, null); + + vfs = new mujoco.MjVFS(); + vfs.addBuffer(objFilename, new TextEncoder().encode(cube1)); + binaryModel = mujoco.MjModel.mj_loadBinary(mjbFilename, vfs); + assertExists(binaryModel); + + expect(mujoco.mj_sizeModel(binaryModel)) + .toEqual(mujoco.mj_sizeModel(model)); + expect(binaryModel.nbody).toEqual(model!.nbody); + expect(binaryModel.nq).toEqual(model!.nq); + expect(binaryModel.nv).toEqual(model!.nv); + expect(binaryModel.njnt).toEqual(model!.njnt); + expect(binaryModel.nmesh).toEqual(model!.nmesh); + } finally { + model?.delete(); + binaryModel?.delete(); + vfs?.delete(); + unlinkXMLFile(xmlFilename); + unlinkXMLFile(objFilename); + unlinkXMLFile(mjbFilename); + } + }); + + // Corresponds to bindings_test.py:test_mj_saveModel + it('should save and load a model from binary', () => { + const mjbFilename = '/tmp/saved_model.mjb'; + let binaryModel: MjModel|null = null; + let vfs: MjVFS|null = null; + try { + mujoco.mj_saveModel(model!, mjbFilename, null); + const bufSize = mujoco.mj_sizeModel(model!); + + vfs = new mujoco.MjVFS(); + binaryModel = mujoco.MjModel.mj_loadBinary(mjbFilename, vfs); + assertExists(binaryModel); + + expect(mujoco.mj_sizeModel(binaryModel)).toEqual(bufSize); + expect(binaryModel.nbody).toEqual(model!.nbody); + expect(binaryModel.nq).toEqual(model!.nq); + expect(binaryModel.nv).toEqual(model!.nv); + expect(binaryModel.njnt).toEqual(model!.njnt); + expect(binaryModel.nmesh).toEqual(model!.nmesh); + } finally { + binaryModel?.delete(); + vfs?.delete(); + unlinkXMLFile(mjbFilename); + } + }); });