From 6fb43947019712eeb41f40dc5e496ab4846dadd6 Mon Sep 17 00:00:00 2001 From: Matias Manevi Date: Fri, 7 Nov 2025 06:11:14 -0800 Subject: [PATCH] Refactor array initialization for `mjData` and `mjvScene` in MuJoCo WASM bindings. PiperOrigin-RevId: 829406147 Change-Id: Ib93c49e938086fb277e4959ade1f25e6ced5c70d --- wasm/codegen/generated/bindings.cc | 91 +++++++----------------------- wasm/codegen/templates/bindings.cc | 91 +++++++----------------------- 2 files changed, 42 insertions(+), 140 deletions(-) diff --git a/wasm/codegen/generated/bindings.cc b/wasm/codegen/generated/bindings.cc index 5418c534..11336581 100644 --- a/wasm/codegen/generated/bindings.cc +++ b/wasm/codegen/generated/bindings.cc @@ -5698,9 +5698,6 @@ struct MjData { MjData(MjModel *m); explicit MjData(const MjModel &, const MjData &); ~MjData(); - std::vector InitSolverArray(); - std::vector InitTimerArray(); - std::vector InitWarningArray(); std::vector contact() const; std::unique_ptr copy(); mjData* get() const; @@ -6381,8 +6378,6 @@ struct MjvScene { ~MjvScene(); std::unique_ptr copy(); int GetSumFlexFaces() const; - std::vector InitLightsArray(); - std::vector InitCameraArray(); mjvScene* get() const; void set(mjvScene* ptr); @@ -8362,6 +8357,15 @@ void MjsDefault::set(mjsDefault* ptr) { ptr_ = ptr; } +template +std::vector InitWrapperArray(ArrayType* array, SizeType size) { + std::vector result; + result.reserve(size); + for (int i = 0; i < size; ++i) { + result.emplace_back(&array[i]); + } + return result; +} // =============== MjModel =============== // MjModel::MjModel(mjModel *m) @@ -8385,17 +8389,17 @@ MjData::MjData(MjModel *m) { model = m->get(); ptr_ = mj_makeData(model); if (ptr_) { - solver = InitSolverArray(); - timer = InitTimerArray(); - warning = InitWarningArray(); + solver = InitWrapperArray(get()->solver, mjNSOLVER * mjNISLAND); + timer = InitWrapperArray(get()->timer, mjNTIMER); + warning = InitWrapperArray(get()->warning, mjNWARNING); } } MjData::MjData(const MjModel &model, const MjData &other) : ptr_(mj_copyData(nullptr, model.get(), other.get())), model(model.get()) { if (ptr_) { - solver = InitSolverArray(); - timer = InitTimerArray(); - warning = InitWarningArray(); + solver = InitWrapperArray(get()->solver, mjNSOLVER * mjNISLAND); + timer = InitWrapperArray(get()->timer, mjNTIMER); + warning = InitWrapperArray(get()->warning, mjNWARNING); } } MjData::~MjData() { @@ -8406,38 +8410,8 @@ MjData::~MjData() { mjData* MjData::get() const { return ptr_; } void MjData::set(mjData *ptr) { ptr_ = ptr; } -std::vector MjData::InitSolverArray() { - std::vector arr; - arr.reserve(mjNSOLVER * mjNISLAND); - for (int i = 0; i < mjNSOLVER * mjNISLAND; i++) { - arr.emplace_back(&get()->solver[i]); - } - return arr; -} -std::vector MjData::InitTimerArray() { - std::vector arr; - arr.reserve(mjNTIMER); - for (int i = 0; i < mjNTIMER; i++) { - arr.emplace_back(&get()->timer[i]); - } - return arr; -} -std::vector -MjData::InitWarningArray() { - std::vector arr; - arr.reserve(mjNWARNING); - for (int i = 0; i < mjNWARNING; i++) { - arr.emplace_back(&get()->warning[i]); - } - return arr; -} std::vector MjData::contact() const { - std::vector contacts; - contacts.reserve(get()->ncon); - for (int i = 0; i < get()->ncon; ++i) { - contacts.emplace_back(&get()->contact[i]); - } - return contacts; + return InitWrapperArray(get()->contact, get()->ncon); } MjvScene::MjvScene() { @@ -8445,8 +8419,8 @@ MjvScene::MjvScene() { ptr_ = new mjvScene; mjv_defaultScene(ptr_); mjv_makeScene(nullptr, ptr_, 0); - lights = InitLightsArray(); - camera = InitCameraArray(); + lights = InitWrapperArray(ptr_->lights, mjMAXLIGHT); + camera = InitWrapperArray(ptr_->camera, 2); }; MjvScene::MjvScene(MjModel *m, int maxgeom) { @@ -8455,8 +8429,8 @@ MjvScene::MjvScene(MjModel *m, int maxgeom) { ptr_ = new mjvScene; mjv_defaultScene(ptr_); mjv_makeScene(model, ptr_, maxgeom); - lights = InitLightsArray(); - camera = InitCameraArray(); + lights = InitWrapperArray(ptr_->lights, mjMAXLIGHT); + camera = InitWrapperArray(ptr_->camera, 2); }; MjvScene::~MjvScene() { if (owned_ && ptr_) { @@ -8502,31 +8476,8 @@ int MjvScene::GetSumFlexFaces() const { return nflexface; } -std::vector MjvScene::InitLightsArray() { - std::vector arr; - arr.reserve(mjMAXLIGHT); - for (int i = 0; i < mjMAXLIGHT; i++) { - arr.emplace_back(&ptr_->lights[i]); - } - return arr; -} - -std::vector MjvScene::InitCameraArray() { - std::vector arr; - arr.reserve(2); - for (int i = 0; i < 2; i++) { - arr.emplace_back(&ptr_->camera[i]); - } - return arr; -} - std::vector MjvScene::geoms() const { - std::vector geoms; - geoms.reserve(ptr_->ngeom); - for (int i = 0; i < ptr_->ngeom; ++i) { - geoms.emplace_back(&ptr_->geoms[i]); - } - return geoms; + return InitWrapperArray(ptr_->geoms, ptr_->ngeom); } MjSpec::MjSpec() diff --git a/wasm/codegen/templates/bindings.cc b/wasm/codegen/templates/bindings.cc index 7d23b464..ba7edf68 100644 --- a/wasm/codegen/templates/bindings.cc +++ b/wasm/codegen/templates/bindings.cc @@ -94,9 +94,6 @@ struct MjData { MjData(MjModel *m); explicit MjData(const MjModel &, const MjData &); ~MjData(); - std::vector InitSolverArray(); - std::vector InitTimerArray(); - std::vector InitWarningArray(); std::vector contact() const; std::unique_ptr copy(); mjData* get() const; @@ -120,8 +117,6 @@ struct MjvScene { ~MjvScene(); std::unique_ptr copy(); int GetSumFlexFaces() const; - std::vector InitLightsArray(); - std::vector InitCameraArray(); mjvScene* get() const; void set(mjvScene* ptr); @@ -352,6 +347,15 @@ EMSCRIPTEN_BINDINGS(mujoco_enums) { // STRUCTS // {{ AUTOGENNED_STRUCTS_SOURCE }} +template +std::vector InitWrapperArray(ArrayType* array, SizeType size) { + std::vector result; + result.reserve(size); + for (int i = 0; i < size; ++i) { + result.emplace_back(&array[i]); + } + return result; +} // =============== MjModel =============== // MjModel::MjModel(mjModel *m) @@ -375,17 +379,17 @@ MjData::MjData(MjModel *m) { model = m->get(); ptr_ = mj_makeData(model); if (ptr_) { - solver = InitSolverArray(); - timer = InitTimerArray(); - warning = InitWarningArray(); + solver = InitWrapperArray(get()->solver, mjNSOLVER * mjNISLAND); + timer = InitWrapperArray(get()->timer, mjNTIMER); + warning = InitWrapperArray(get()->warning, mjNWARNING); } } MjData::MjData(const MjModel &model, const MjData &other) : ptr_(mj_copyData(nullptr, model.get(), other.get())), model(model.get()) { if (ptr_) { - solver = InitSolverArray(); - timer = InitTimerArray(); - warning = InitWarningArray(); + solver = InitWrapperArray(get()->solver, mjNSOLVER * mjNISLAND); + timer = InitWrapperArray(get()->timer, mjNTIMER); + warning = InitWrapperArray(get()->warning, mjNWARNING); } } MjData::~MjData() { @@ -396,38 +400,8 @@ MjData::~MjData() { mjData* MjData::get() const { return ptr_; } void MjData::set(mjData *ptr) { ptr_ = ptr; } -std::vector MjData::InitSolverArray() { - std::vector arr; - arr.reserve(mjNSOLVER * mjNISLAND); - for (int i = 0; i < mjNSOLVER * mjNISLAND; i++) { - arr.emplace_back(&get()->solver[i]); - } - return arr; -} -std::vector MjData::InitTimerArray() { - std::vector arr; - arr.reserve(mjNTIMER); - for (int i = 0; i < mjNTIMER; i++) { - arr.emplace_back(&get()->timer[i]); - } - return arr; -} -std::vector -MjData::InitWarningArray() { - std::vector arr; - arr.reserve(mjNWARNING); - for (int i = 0; i < mjNWARNING; i++) { - arr.emplace_back(&get()->warning[i]); - } - return arr; -} std::vector MjData::contact() const { - std::vector contacts; - contacts.reserve(get()->ncon); - for (int i = 0; i < get()->ncon; ++i) { - contacts.emplace_back(&get()->contact[i]); - } - return contacts; + return InitWrapperArray(get()->contact, get()->ncon); } MjvScene::MjvScene() { @@ -435,8 +409,8 @@ MjvScene::MjvScene() { ptr_ = new mjvScene; mjv_defaultScene(ptr_); mjv_makeScene(nullptr, ptr_, 0); - lights = InitLightsArray(); - camera = InitCameraArray(); + lights = InitWrapperArray(ptr_->lights, mjMAXLIGHT); + camera = InitWrapperArray(ptr_->camera, 2); }; MjvScene::MjvScene(MjModel *m, int maxgeom) { @@ -445,8 +419,8 @@ MjvScene::MjvScene(MjModel *m, int maxgeom) { ptr_ = new mjvScene; mjv_defaultScene(ptr_); mjv_makeScene(model, ptr_, maxgeom); - lights = InitLightsArray(); - camera = InitCameraArray(); + lights = InitWrapperArray(ptr_->lights, mjMAXLIGHT); + camera = InitWrapperArray(ptr_->camera, 2); }; MjvScene::~MjvScene() { if (owned_ && ptr_) { @@ -492,31 +466,8 @@ int MjvScene::GetSumFlexFaces() const { return nflexface; } -std::vector MjvScene::InitLightsArray() { - std::vector arr; - arr.reserve(mjMAXLIGHT); - for (int i = 0; i < mjMAXLIGHT; i++) { - arr.emplace_back(&ptr_->lights[i]); - } - return arr; -} - -std::vector MjvScene::InitCameraArray() { - std::vector arr; - arr.reserve(2); - for (int i = 0; i < 2; i++) { - arr.emplace_back(&ptr_->camera[i]); - } - return arr; -} - std::vector MjvScene::geoms() const { - std::vector geoms; - geoms.reserve(ptr_->ngeom); - for (int i = 0; i < ptr_->ngeom; ++i) { - geoms.emplace_back(&ptr_->geoms[i]); - } - return geoms; + return InitWrapperArray(ptr_->geoms, ptr_->ngeom); } MjSpec::MjSpec()