Add code generation for mjModel named accessors.
Most of it is done leveraging existing X Macros (indexers.h, indexer_xmacro.h) used for python bindings already. PiperOrigin-RevId: 856205245 Change-Id: Ic72443999065040af3042bbf3f9f68c5b70d309d
This commit is contained in:
committed by
Copybara-Service
parent
cc4c2bc398
commit
c996f288e1
@@ -28,6 +28,7 @@
|
||||
#include <memory>
|
||||
#include <optional> // NOLINT
|
||||
#include <string> // NOLINT
|
||||
#include <string_view>
|
||||
#include <vector>
|
||||
|
||||
#include <mujoco/mjmodel.h>
|
||||
@@ -36,6 +37,7 @@
|
||||
#include <mujoco/mujoco.h>
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "wasm/unpack.h"
|
||||
#include "python/mujoco/indexer_xmacro.h"
|
||||
|
||||
namespace mujoco::wasm {
|
||||
|
||||
@@ -45,9 +47,39 @@ using emscripten::val;
|
||||
using emscripten::return_value_policy::reference;
|
||||
using emscripten::return_value_policy::take_ownership;
|
||||
|
||||
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberOrString);
|
||||
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray);
|
||||
EMSCRIPTEN_DECLARE_VAL_TYPE(String);
|
||||
|
||||
// Macro to define accessors for different MuJoCo object types within mjModel.
|
||||
// Each line calls X_ACCESSOR with the following arguments:
|
||||
// 1. NAME: The object type name in uppercase (e.g., ACTUATOR).
|
||||
// 2. Name: The object type name in CamelCase (e.g., Actuator).
|
||||
// 3. OBJTYPE: The corresponding mjOBJ_* enum value (e.g., mjOBJ_ACTUATOR).
|
||||
// 4. field: The name of the array field in mjModel (e.g., actuator).
|
||||
// 5. nfield: The name of the count field in mjModel (e.g., nu).
|
||||
#define MJMODEL_ACCESSORS \
|
||||
X_ACCESSOR( ACTUATOR, Actuator, mjOBJ_ACTUATOR, actuator, nu ) \
|
||||
X_ACCESSOR( BODY, Body, mjOBJ_BODY, body, nbody ) \
|
||||
X_ACCESSOR( CAMERA, Camera, mjOBJ_CAMERA, cam, ncam ) \
|
||||
X_ACCESSOR( EQUALITY, Equality, mjOBJ_EQUALITY, eq, neq ) \
|
||||
X_ACCESSOR( EXCLUDE, Exclude, mjOBJ_EXCLUDE, exclude, nexclude ) \
|
||||
X_ACCESSOR( GEOM, Geom, mjOBJ_GEOM, geom, ngeom ) \
|
||||
X_ACCESSOR( HFIELD, Hfield, mjOBJ_HFIELD, hfield, nhfield ) \
|
||||
X_ACCESSOR( JOINT, Joint, mjOBJ_JOINT, jnt, njnt ) \
|
||||
X_ACCESSOR( LIGHT, Light, mjOBJ_LIGHT, light, nlight ) \
|
||||
X_ACCESSOR( MATERIAL, Material, mjOBJ_MATERIAL, mat, nmat ) \
|
||||
X_ACCESSOR( MESH, Mesh, mjOBJ_MESH, mesh, nmesh ) \
|
||||
X_ACCESSOR( NUMERIC, Numeric, mjOBJ_NUMERIC, numeric, nnumeric ) \
|
||||
X_ACCESSOR( PAIR, Pair, mjOBJ_PAIR, pair, npair ) \
|
||||
X_ACCESSOR( SENSOR, Sensor, mjOBJ_SENSOR, sensor, nsensor ) \
|
||||
X_ACCESSOR( SITE, Site, mjOBJ_SITE, site, nsite ) \
|
||||
X_ACCESSOR( SKIN, Skin, mjOBJ_SKIN, skin, nskin ) \
|
||||
X_ACCESSOR( TENDON, Tendon, mjOBJ_TENDON, tendon, ntendon ) \
|
||||
X_ACCESSOR( TEXTURE, Texture, mjOBJ_TEXTURE, tex, ntex ) \
|
||||
X_ACCESSOR( TUPLE, Tuple, mjOBJ_TUPLE, tuple, ntuple ) \
|
||||
X_ACCESSOR( KEYFRAME, Keyframe, mjOBJ_KEY, key, nkey )
|
||||
|
||||
// Raises an error if the given val is null or undefined.
|
||||
// A macro is used so that the error contains the name of the variable.
|
||||
// TODO(matijak): Remove this when we can handle strings using UNPACK_STRING?
|
||||
@@ -117,6 +149,84 @@ using mjVisualQuality = decltype(::mjVisual::quality);
|
||||
using mjVisualRgba = decltype(::mjVisual::rgba);
|
||||
using mjVisualScale = decltype(::mjVisual::scale);
|
||||
|
||||
#undef MJ_M
|
||||
#define MJ_M(n) model_->n
|
||||
// The X macro expands to a member function within an MjModel...Accessor struct.
|
||||
// This function returns an emscripten::val, typically a typed memory view,
|
||||
// providing access to array data within the underlying mjModel. The size and
|
||||
// offset of the memory view are determined by the arguments:
|
||||
// - type: The C++ type of the array elements (e.g., mjtNum, int).
|
||||
// - prefix: The prefix of the field name in mjModel (e.g., jnt_, geom_).
|
||||
// - var: The name of the field being accessed (e.g., qposadr, size).
|
||||
// - dim0: Used to determine special indexing logic for dynamically sized arrays
|
||||
// like joints (nq, nv), hfields, textures, numerics, and tuples.
|
||||
// - dim1: The fixed dimension of the array if not dynamically sized. If "1",
|
||||
// a single element is returned. Otherwise, it's used as the stride
|
||||
// for the typed memory view.
|
||||
#define X(type, prefix, var, dim0, dim1) \
|
||||
auto var() const { \
|
||||
if constexpr (std::string_view(#dim0) == "nq") { \
|
||||
int start = model_->jnt_qposadr[id_]; \
|
||||
int end = (id_ < model_->njnt - 1) ? model_->jnt_qposadr[id_ + 1] : model_->nq; \
|
||||
return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "nv") { \
|
||||
int start = model_->jnt_dofadr[id_]; \
|
||||
int end = (id_ < model_->njnt - 1) ? model_->jnt_dofadr[id_ + 1] : model_->nv; \
|
||||
return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "nhfielddata") { \
|
||||
int start = model_->hfield_adr[id_]; \
|
||||
int count = model_->hfield_nrow[id_] * model_->hfield_ncol[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->hfield_data + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "ntexdata") { \
|
||||
int start = model_->tex_adr[id_]; \
|
||||
int count = model_->tex_height[id_] * model_->tex_width[id_] * model_->tex_nchannel[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->tex_data + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "nnumericdata") { \
|
||||
int start = model_->numeric_adr[id_]; \
|
||||
int count = model_->numeric_size[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->numeric_data + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "ntupledata") { \
|
||||
int start = model_->tuple_adr[id_]; \
|
||||
int count = model_->tuple_size[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->prefix##var + start)); \
|
||||
} else { \
|
||||
if constexpr (std::string_view(#dim1) == "1") { \
|
||||
return model_->prefix##var[id_]; \
|
||||
} else { \
|
||||
return emscripten::val( \
|
||||
emscripten::typed_memory_view(dim1, model_->prefix##var + id_ * dim1)); \
|
||||
} \
|
||||
} \
|
||||
}
|
||||
|
||||
// Expands to a struct definition for each object type in MJMODEL_ACCESSORS.
|
||||
// Each struct, named `MjModel{Name}Accessor`, provides:
|
||||
// - A constructor taking an `mjModel*` and an integer `id`.
|
||||
// - An `id()` method to get the object's index.
|
||||
// - A `name()` method to get the object's name using `mj_id2name`.
|
||||
// - Member functions generated by the `MJMODEL_##NAME` macro, which in turn
|
||||
// uses the `X` macro to define accessors for fields within `mjModel`.
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \
|
||||
struct MjModel##Name##Accessor { \
|
||||
MjModel##Name##Accessor(mjModel* model, int id) : model_(model), id_(id) {} \
|
||||
\
|
||||
int id() const { return id_; } \
|
||||
std::string name() const { \
|
||||
return mj_id2name(model_, OBJTYPE, id_); \
|
||||
} \
|
||||
\
|
||||
MJMODEL_##NAME \
|
||||
\
|
||||
private: \
|
||||
mjModel* model_; \
|
||||
int id_; \
|
||||
};
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X_ACCESSOR
|
||||
|
||||
#undef X
|
||||
#undef MJ_M
|
||||
|
||||
struct MjContact {
|
||||
~MjContact();
|
||||
MjContact();
|
||||
@@ -5143,6 +5253,29 @@ struct MjModel {
|
||||
void set_signature(uint64_t value) {
|
||||
ptr_->signature = value;
|
||||
}
|
||||
// Generates functions to return accessor classes.
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, accessor_name, nfield) \
|
||||
MjModel##Name##Accessor accessor_name(const NumberOrString& val) const { \
|
||||
if (val.isString()) { \
|
||||
int id = mj_name2id(ptr_, OBJTYPE, val.as<std::string>().c_str()); \
|
||||
if (id == -1) { \
|
||||
mju_error("Invalid name, MjModel." #accessor_name " not found"); \
|
||||
} \
|
||||
return MjModel##Name##Accessor(ptr_, id); \
|
||||
} else if (val.isNumber()) { \
|
||||
int id = val.as<int>(); \
|
||||
if (id < 0 || id >= ptr_->nfield) { \
|
||||
mju_error_i("Invalid id %d for MjModel." #accessor_name, id); \
|
||||
} \
|
||||
return MjModel##Name##Accessor(ptr_, id); \
|
||||
} else { \
|
||||
mju_error(#accessor_name "() argument must be a string or number"); \
|
||||
return MjModel##Name##Accessor(nullptr, 0); \
|
||||
} \
|
||||
}
|
||||
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X_ACCESSOR
|
||||
|
||||
private:
|
||||
mjModel* ptr_;
|
||||
@@ -10784,6 +10917,20 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
|
||||
.value("mjWRAP_SPHERE", mjWRAP_SPHERE)
|
||||
.value("mjWRAP_CYLINDER", mjWRAP_CYLINDER);
|
||||
|
||||
// Bindings for the MjModel accessor classes.
|
||||
#define X(type, prefix, var, dim0, dim1) .property(#var, &Accessor::var)
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \
|
||||
{ \
|
||||
using Accessor = MjModel##Name##Accessor; \
|
||||
emscripten::class_<Accessor>("MjModel" #Name "Accessor") \
|
||||
.property("id", &Accessor::id) \
|
||||
.property("name", &Accessor::name) \
|
||||
MJMODEL_##NAME; \
|
||||
}
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X
|
||||
#undef X_ACCESSOR
|
||||
|
||||
emscripten::class_<MjContact>("MjContact")
|
||||
.constructor<>()
|
||||
.function("copy", &MjContact::copy, take_ownership())
|
||||
@@ -11018,6 +11165,11 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
|
||||
emscripten::class_<MjModel>("MjModel")
|
||||
.class_function("mj_loadXML", &mj_loadXML_wrapper, take_ownership())
|
||||
.constructor<const MjModel &>()
|
||||
// Binds the functions on MjModel that return accessors.
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, field_name, nfield) \
|
||||
.function(#field_name, &MjModel::field_name)
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X_ACCESSOR
|
||||
.property("B_colind", &MjModel::B_colind)
|
||||
.property("B_rowadr", &MjModel::B_rowadr)
|
||||
.property("B_rownnz", &MjModel::B_rownnz)
|
||||
@@ -12795,6 +12947,9 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
|
||||
emscripten::register_vector<MjvGeom>("MjvGeomVec");
|
||||
|
||||
// register_type() improves type information (val is mapped to any by default)
|
||||
// NumberOrString is used in functions returning accessors, allowing users to
|
||||
// get an accessor by name (string) or id (number).
|
||||
emscripten::register_type<NumberOrString>("number|string");
|
||||
emscripten::register_type<NumberArray>("number[]");
|
||||
emscripten::register_type<String>("string");
|
||||
|
||||
|
||||
@@ -20,8 +20,8 @@ import math
|
||||
from typing import Tuple, Union, cast
|
||||
|
||||
from introspect import ast_nodes
|
||||
from introspect import structs as introspect_structs
|
||||
from introspect import functions as introspect_functions
|
||||
from introspect import structs as introspect_structs
|
||||
|
||||
from wasm.codegen.generators import code_builder
|
||||
from wasm.codegen.generators import common
|
||||
@@ -376,6 +376,33 @@ def build_struct_header(
|
||||
for line in field.declaration.splitlines():
|
||||
builder.line(line)
|
||||
|
||||
# accessors declarations
|
||||
if w == "MjModel":
|
||||
builder.line("""
|
||||
// Generates functions to return accessor classes.
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, accessor_name, nfield) \\
|
||||
MjModel##Name##Accessor accessor_name(const NumberOrString& val) const { \\
|
||||
if (val.isString()) { \\
|
||||
int id = mj_name2id(ptr_, OBJTYPE, val.as<std::string>().c_str()); \\
|
||||
if (id == -1) { \\
|
||||
mju_error("Invalid name, MjModel." #accessor_name " not found"); \\
|
||||
} \\
|
||||
return MjModel##Name##Accessor(ptr_, id); \\
|
||||
} else if (val.isNumber()) { \\
|
||||
int id = val.as<int>(); \\
|
||||
if (id < 0 || id >= ptr_->nfield) { \\
|
||||
mju_error_i("Invalid id %d for MjModel." #accessor_name, id); \\
|
||||
} \\
|
||||
return MjModel##Name##Accessor(ptr_, id); \\
|
||||
} else { \\
|
||||
mju_error(#accessor_name "() argument must be a string or number"); \\
|
||||
return MjModel##Name##Accessor(nullptr, 0); \\
|
||||
} \\
|
||||
}
|
||||
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X_ACCESSOR""".lstrip())
|
||||
|
||||
# define private struct members
|
||||
builder.private()
|
||||
builder.line(f"{s}* ptr_;")
|
||||
@@ -481,11 +508,17 @@ def _build_struct_bindings(
|
||||
builder.line(".constructor<MjModel *>()")
|
||||
builder.line(".constructor<const MjModel &, const MjData &>()")
|
||||
elif w == "MjModel":
|
||||
w = common.wrapped_function_name(
|
||||
fn = common.wrapped_function_name(
|
||||
introspect_functions.FUNCTIONS["mj_loadXML"]
|
||||
)
|
||||
builder.line(f'.class_function("mj_loadXML", &{w}, take_ownership())')
|
||||
builder.line(f'.class_function("mj_loadXML", &{fn}, take_ownership())')
|
||||
builder.line(".constructor<const MjModel &>()")
|
||||
builder.line("""
|
||||
// Binds the functions on MjModel that return accessors.
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, field_name, nfield) \\
|
||||
.function(#field_name, &MjModel::field_name)
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X_ACCESSOR""".lstrip())
|
||||
elif w == "MjSpec":
|
||||
builder.line(".constructor<const MjSpec &>()")
|
||||
elif w == "MjvScene":
|
||||
|
||||
@@ -28,6 +28,7 @@
|
||||
#include <memory>
|
||||
#include <optional> // NOLINT
|
||||
#include <string> // NOLINT
|
||||
#include <string_view>
|
||||
#include <vector>
|
||||
|
||||
#include <mujoco/mjmodel.h>
|
||||
@@ -36,6 +37,7 @@
|
||||
#include <mujoco/mujoco.h>
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "wasm/unpack.h"
|
||||
#include "python/mujoco/indexer_xmacro.h"
|
||||
|
||||
namespace mujoco::wasm {
|
||||
|
||||
@@ -45,9 +47,39 @@ using emscripten::val;
|
||||
using emscripten::return_value_policy::reference;
|
||||
using emscripten::return_value_policy::take_ownership;
|
||||
|
||||
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberOrString);
|
||||
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray);
|
||||
EMSCRIPTEN_DECLARE_VAL_TYPE(String);
|
||||
|
||||
// Macro to define accessors for different MuJoCo object types within mjModel.
|
||||
// Each line calls X_ACCESSOR with the following arguments:
|
||||
// 1. NAME: The object type name in uppercase (e.g., ACTUATOR).
|
||||
// 2. Name: The object type name in CamelCase (e.g., Actuator).
|
||||
// 3. OBJTYPE: The corresponding mjOBJ_* enum value (e.g., mjOBJ_ACTUATOR).
|
||||
// 4. field: The name of the array field in mjModel (e.g., actuator).
|
||||
// 5. nfield: The name of the count field in mjModel (e.g., nu).
|
||||
#define MJMODEL_ACCESSORS \
|
||||
X_ACCESSOR( ACTUATOR, Actuator, mjOBJ_ACTUATOR, actuator, nu ) \
|
||||
X_ACCESSOR( BODY, Body, mjOBJ_BODY, body, nbody ) \
|
||||
X_ACCESSOR( CAMERA, Camera, mjOBJ_CAMERA, cam, ncam ) \
|
||||
X_ACCESSOR( EQUALITY, Equality, mjOBJ_EQUALITY, eq, neq ) \
|
||||
X_ACCESSOR( EXCLUDE, Exclude, mjOBJ_EXCLUDE, exclude, nexclude ) \
|
||||
X_ACCESSOR( GEOM, Geom, mjOBJ_GEOM, geom, ngeom ) \
|
||||
X_ACCESSOR( HFIELD, Hfield, mjOBJ_HFIELD, hfield, nhfield ) \
|
||||
X_ACCESSOR( JOINT, Joint, mjOBJ_JOINT, jnt, njnt ) \
|
||||
X_ACCESSOR( LIGHT, Light, mjOBJ_LIGHT, light, nlight ) \
|
||||
X_ACCESSOR( MATERIAL, Material, mjOBJ_MATERIAL, mat, nmat ) \
|
||||
X_ACCESSOR( MESH, Mesh, mjOBJ_MESH, mesh, nmesh ) \
|
||||
X_ACCESSOR( NUMERIC, Numeric, mjOBJ_NUMERIC, numeric, nnumeric ) \
|
||||
X_ACCESSOR( PAIR, Pair, mjOBJ_PAIR, pair, npair ) \
|
||||
X_ACCESSOR( SENSOR, Sensor, mjOBJ_SENSOR, sensor, nsensor ) \
|
||||
X_ACCESSOR( SITE, Site, mjOBJ_SITE, site, nsite ) \
|
||||
X_ACCESSOR( SKIN, Skin, mjOBJ_SKIN, skin, nskin ) \
|
||||
X_ACCESSOR( TENDON, Tendon, mjOBJ_TENDON, tendon, ntendon ) \
|
||||
X_ACCESSOR( TEXTURE, Texture, mjOBJ_TEXTURE, tex, ntex ) \
|
||||
X_ACCESSOR( TUPLE, Tuple, mjOBJ_TUPLE, tuple, ntuple ) \
|
||||
X_ACCESSOR( KEYFRAME, Keyframe, mjOBJ_KEY, key, nkey )
|
||||
|
||||
// Raises an error if the given val is null or undefined.
|
||||
// A macro is used so that the error contains the name of the variable.
|
||||
// TODO(matijak): Remove this when we can handle strings using UNPACK_STRING?
|
||||
@@ -112,6 +144,84 @@ val get_mjRNDSTRING() { return MakeValArray3(mjRNDSTRING); }
|
||||
|
||||
// {{ ANONYMOUS_STRUCT_TYPEDEFS }}
|
||||
|
||||
#undef MJ_M
|
||||
#define MJ_M(n) model_->n
|
||||
// The X macro expands to a member function within an MjModel...Accessor struct.
|
||||
// This function returns an emscripten::val, typically a typed memory view,
|
||||
// providing access to array data within the underlying mjModel. The size and
|
||||
// offset of the memory view are determined by the arguments:
|
||||
// - type: The C++ type of the array elements (e.g., mjtNum, int).
|
||||
// - prefix: The prefix of the field name in mjModel (e.g., jnt_, geom_).
|
||||
// - var: The name of the field being accessed (e.g., qposadr, size).
|
||||
// - dim0: Used to determine special indexing logic for dynamically sized arrays
|
||||
// like joints (nq, nv), hfields, textures, numerics, and tuples.
|
||||
// - dim1: The fixed dimension of the array if not dynamically sized. If "1",
|
||||
// a single element is returned. Otherwise, it's used as the stride
|
||||
// for the typed memory view.
|
||||
#define X(type, prefix, var, dim0, dim1) \
|
||||
auto var() const { \
|
||||
if constexpr (std::string_view(#dim0) == "nq") { \
|
||||
int start = model_->jnt_qposadr[id_]; \
|
||||
int end = (id_ < model_->njnt - 1) ? model_->jnt_qposadr[id_ + 1] : model_->nq; \
|
||||
return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "nv") { \
|
||||
int start = model_->jnt_dofadr[id_]; \
|
||||
int end = (id_ < model_->njnt - 1) ? model_->jnt_dofadr[id_ + 1] : model_->nv; \
|
||||
return emscripten::val(emscripten::typed_memory_view(end - start, model_->prefix##var + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "nhfielddata") { \
|
||||
int start = model_->hfield_adr[id_]; \
|
||||
int count = model_->hfield_nrow[id_] * model_->hfield_ncol[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->hfield_data + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "ntexdata") { \
|
||||
int start = model_->tex_adr[id_]; \
|
||||
int count = model_->tex_height[id_] * model_->tex_width[id_] * model_->tex_nchannel[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->tex_data + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "nnumericdata") { \
|
||||
int start = model_->numeric_adr[id_]; \
|
||||
int count = model_->numeric_size[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->numeric_data + start)); \
|
||||
} else if constexpr (std::string_view(#dim0) == "ntupledata") { \
|
||||
int start = model_->tuple_adr[id_]; \
|
||||
int count = model_->tuple_size[id_]; \
|
||||
return emscripten::val(emscripten::typed_memory_view(count, model_->prefix##var + start)); \
|
||||
} else { \
|
||||
if constexpr (std::string_view(#dim1) == "1") { \
|
||||
return model_->prefix##var[id_]; \
|
||||
} else { \
|
||||
return emscripten::val( \
|
||||
emscripten::typed_memory_view(dim1, model_->prefix##var + id_ * dim1)); \
|
||||
} \
|
||||
} \
|
||||
}
|
||||
|
||||
// Expands to a struct definition for each object type in MJMODEL_ACCESSORS.
|
||||
// Each struct, named `MjModel{Name}Accessor`, provides:
|
||||
// - A constructor taking an `mjModel*` and an integer `id`.
|
||||
// - An `id()` method to get the object's index.
|
||||
// - A `name()` method to get the object's name using `mj_id2name`.
|
||||
// - Member functions generated by the `MJMODEL_##NAME` macro, which in turn
|
||||
// uses the `X` macro to define accessors for fields within `mjModel`.
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \
|
||||
struct MjModel##Name##Accessor { \
|
||||
MjModel##Name##Accessor(mjModel* model, int id) : model_(model), id_(id) {} \
|
||||
\
|
||||
int id() const { return id_; } \
|
||||
std::string name() const { \
|
||||
return mj_id2name(model_, OBJTYPE, id_); \
|
||||
} \
|
||||
\
|
||||
MJMODEL_##NAME \
|
||||
\
|
||||
private: \
|
||||
mjModel* model_; \
|
||||
int id_; \
|
||||
};
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X_ACCESSOR
|
||||
|
||||
#undef X
|
||||
#undef MJ_M
|
||||
|
||||
// {{ STRUCTS_HEADER }}
|
||||
|
||||
// {{ STRUCTS_SOURCE }}
|
||||
@@ -473,6 +583,20 @@ int mj_setLengthRange_wrapper(const MjModel& m, const MjData& d, int index, cons
|
||||
EMSCRIPTEN_BINDINGS(mujoco_bindings) {
|
||||
// {{ ENUM_BINDINGS }}
|
||||
|
||||
// Bindings for the MjModel accessor classes.
|
||||
#define X(type, prefix, var, dim0, dim1) .property(#var, &Accessor::var)
|
||||
#define X_ACCESSOR(NAME, Name, OBJTYPE, field, nfield) \
|
||||
{ \
|
||||
using Accessor = MjModel##Name##Accessor; \
|
||||
emscripten::class_<Accessor>("MjModel" #Name "Accessor") \
|
||||
.property("id", &Accessor::id) \
|
||||
.property("name", &Accessor::name) \
|
||||
MJMODEL_##NAME; \
|
||||
}
|
||||
MJMODEL_ACCESSORS
|
||||
#undef X
|
||||
#undef X_ACCESSOR
|
||||
|
||||
// {{ STRUCTS_BINDINGS }}
|
||||
|
||||
emscripten::class_<MjVFS>("MjVFS")
|
||||
@@ -527,6 +651,9 @@ EMSCRIPTEN_BINDINGS(mujoco_bindings) {
|
||||
emscripten::register_vector<MjvGeom>("MjvGeomVec");
|
||||
|
||||
// register_type() improves type information (val is mapped to any by default)
|
||||
// NumberOrString is used in functions returning accessors, allowing users to
|
||||
// get an accessor by name (string) or id (number).
|
||||
emscripten::register_type<NumberOrString>("number|string");
|
||||
emscripten::register_type<NumberArray>("number[]");
|
||||
emscripten::register_type<String>("string");
|
||||
|
||||
|
||||
+270
-2
@@ -75,14 +75,14 @@ function norm(arr: number[]): number {
|
||||
return Math.sqrt(arr.reduce((acc, val) => acc + val * val, 0));
|
||||
}
|
||||
|
||||
function expectArraysClose(arr1: TypedArray, arr2: TypedArray, precision = 1) {
|
||||
function expectArraysClose(arr1: any, arr2: TypedArray, precision = 1) {
|
||||
expect(arr1.length).toEqual(arr2.length);
|
||||
for (let i = 0; i < arr1.length; i++) {
|
||||
expect(arr1[i]).toBeCloseTo(arr2[i], precision);
|
||||
}
|
||||
}
|
||||
|
||||
function expectArraysEqual(arr1: TypedArray, arr2: TypedArray) {
|
||||
function expectArraysEqual(arr1: any, arr2: TypedArray) {
|
||||
expect(arr1.length).toEqual(arr2.length);
|
||||
for (let i = 0; i < arr1.length; i++) {
|
||||
expect(arr1[i]).toEqual(arr2[i]);
|
||||
@@ -1872,4 +1872,272 @@ describe('MuJoCo WASM Bindings', () => {
|
||||
model?.delete();
|
||||
}
|
||||
});
|
||||
|
||||
describe('MjModel named access', () => {
|
||||
// Corresponds to
|
||||
// bindings_test.py:test_named_indexing_invalid_names_in_model
|
||||
it('should throw an error for invalid geom names in model', () => {
|
||||
expect(() => {
|
||||
model!.geom('badgeom');
|
||||
}).toThrowError('MuJoCo Error: Invalid name, MjModel.geom not found');
|
||||
});
|
||||
|
||||
// Corresponds to bindings_test.py:test_named_indexing_geom_size
|
||||
it('should correctly access geom size using named indexing', () => {
|
||||
const boxId =
|
||||
mujoco.mj_name2id(model!, mujoco.mjtObj.mjOBJ_GEOM.value, 'mybox');
|
||||
assertExists(boxId);
|
||||
|
||||
// Check that named indexing returns the same object as id indexing.
|
||||
expect(model!.geom('mybox')).toEqual(model!.geom(boxId));
|
||||
expect(model!.geom('mybox')!.size).toEqual(model!.geom(boxId)!.size);
|
||||
expect(model!.geom('mybox')!.size.length).toEqual(3);
|
||||
|
||||
// Test that the indexer is returning a view into the underlying struct.
|
||||
const sizeFromIndexer = model!.geom('mybox')!.size;
|
||||
const originalGeomSize = new Float64Array(model!.geom_size);
|
||||
|
||||
model!.geom_size.set([7, 11, 13], boxId * 3);
|
||||
expectArraysEqual(sizeFromIndexer, new Float64Array([7, 11, 13]));
|
||||
|
||||
model!.geom('mybox')!.size.set([5, 3, 2]);
|
||||
expectArraysEqual(
|
||||
model!.geom_size.slice(boxId * 3, boxId * 3 + 3),
|
||||
new Float64Array([5, 3, 2]));
|
||||
});
|
||||
|
||||
// Corresponds to bindings_test.py:test_named_indexing_geom_quat
|
||||
it('should correctly access geom quat using named indexing', () => {
|
||||
const boxId =
|
||||
mujoco.mj_name2id(model!, mujoco.mjtObj.mjOBJ_GEOM.value, 'mybox');
|
||||
assertExists(boxId);
|
||||
|
||||
// Check that named indexing returns the same object as id indexing.
|
||||
expect(model!.geom('mybox')).toEqual(model!.geom(boxId));
|
||||
expect(model!.geom('mybox')!.quat).toEqual(model!.geom(boxId)!.quat);
|
||||
expect(model!.geom('mybox')!.quat.length).toEqual(4);
|
||||
|
||||
// Test that the indexer is returning a view into the underlying struct.
|
||||
const quatFromIndexer = model!.geom('mybox')!.quat;
|
||||
|
||||
model!.geom_quat.set([5, 10, 15, 20], boxId * 4);
|
||||
expectArraysEqual(quatFromIndexer, new Float64Array([5, 10, 15, 20]));
|
||||
|
||||
model!.geom('mybox')!.quat.set([12, 9, 6, 3]);
|
||||
expectArraysEqual(
|
||||
model!.geom_quat.slice(boxId * 4, boxId * 4 + 4),
|
||||
new Float64Array([12, 9, 6, 3]));
|
||||
});
|
||||
|
||||
it('should support named access for Joints', () => {
|
||||
const xml = `<mujoco model="Box falling">
|
||||
<option viscosity="1"/>
|
||||
<worldbody>
|
||||
<light diffuse=".5 .5 .5" pos="0 0 3" dir="0 0 -1"/>
|
||||
<geom name="MyFloor" type="plane" size="1 1 0.1" rgba=".9 0 0 1" user="5 4 3 2 1"/>
|
||||
<geom name="MyWall" type="plane" size="0.1 1 0.1" rgba="0 1 0 1" user="5 4 3"/>
|
||||
<body pos="0 0 1" name="MyBox">
|
||||
<joint type="free" name="root"/>
|
||||
<geom name="MyBoxGeom" type="box" size=".1 .2 .3" rgba="0 .9 0 1"/>
|
||||
<body pos=".1 0 0" name="MyHingeBody">
|
||||
<joint frictionloss="0.4" type="hinge" name="hinge" axis="0 1 0"/>
|
||||
<geom type="capsule" size=".05 .05" rgba="0 0 .9 1"/>
|
||||
</body>
|
||||
</body>
|
||||
</worldbody>
|
||||
</mujoco>`;
|
||||
const tempXmlFilename = '/tmp/jnt.xml';
|
||||
writeXMLFile(tempXmlFilename, xml);
|
||||
const model = mujoco.MjModel.mj_loadXML(tempXmlFilename);
|
||||
try {
|
||||
assertExists(model);
|
||||
const j0 = model.jnt('root');
|
||||
expect(j0.id).toBe(0);
|
||||
expect(j0.name).toBe('root');
|
||||
expect(j0.type).toBe(mujoco.mjtJoint.mjJNT_FREE.value);
|
||||
expect(j0.qposadr).toBe(0);
|
||||
expect(j0.dofadr).toBe(0);
|
||||
expect(j0.axis).toEqual(new Float64Array([0, 0, 1]));
|
||||
expect(j0.type).toEqual(mujoco.mjtJoint.mjJNT_FREE.value);
|
||||
|
||||
const j0qposStart = model.jnt_qposadr[j0.id];
|
||||
const j0qposEnd =
|
||||
(j0.id < model.njnt - 1) ? model.jnt_qposadr[j0.id + 1] : model.nq;
|
||||
const expectedJ0qpos = model.qpos0.slice(j0qposStart, j0qposEnd);
|
||||
expectArraysEqual(
|
||||
expectedJ0qpos, new Float64Array([0, 0, 1, 1, 0, 0, 0]));
|
||||
expectArraysEqual(j0.qpos0, expectedJ0qpos);
|
||||
|
||||
const j0dofStart = model.jnt_dofadr[j0.id];
|
||||
const j0dofEnd =
|
||||
(j0.id < model.njnt - 1) ? model.jnt_dofadr[j0.id + 1] : model.nv;
|
||||
const expectedJ0dof = model.dof_bodyid.slice(j0dofStart, j0dofEnd);
|
||||
expectArraysEqual(expectedJ0dof, new Int32Array([1, 1, 1, 1, 1, 1]));
|
||||
expectArraysEqual(j0.bodyid, expectedJ0dof);
|
||||
|
||||
|
||||
const j1 = model.jnt('hinge');
|
||||
expect(j1.id).toBe(1);
|
||||
expect(j1.name).toBe('hinge');
|
||||
expect(j1.type).toBe(mujoco.mjtJoint.mjJNT_HINGE.value);
|
||||
expect(j1.qposadr).toBe(7);
|
||||
expect(j1.dofadr).toBe(6);
|
||||
expect(j1.axis).toEqual(new Float64Array([0, 1, 0]));
|
||||
expect(j1.frictionloss).toEqual(new Float64Array([0.4]));
|
||||
|
||||
const j1qposStart = model.jnt_qposadr[j1.id];
|
||||
const j1qposEnd =
|
||||
(j1.id < model.njnt - 1) ? model.jnt_qposadr[j1.id + 1] : model.nq;
|
||||
const expectedJ1qpos = model.qpos0.slice(j1qposStart, j1qposEnd);
|
||||
expectArraysEqual(expectedJ1qpos, new Float64Array([0]));
|
||||
expectArraysEqual(j1.qpos0, expectedJ1qpos);
|
||||
|
||||
const j1dofStart = model.jnt_dofadr[j1.id];
|
||||
const j1dofEnd =
|
||||
(j1.id < model.njnt - 1) ? model.jnt_dofadr[j1.id + 1] : model.nv;
|
||||
const expectedJ1dof = model.dof_bodyid.slice(j1dofStart, j1dofEnd);
|
||||
expectArraysEqual(expectedJ1dof, new Int32Array([2]));
|
||||
expectArraysEqual(j1.bodyid, new Int32Array([2]));
|
||||
} finally {
|
||||
model?.delete();
|
||||
unlinkXMLFile(tempXmlFilename);
|
||||
}
|
||||
});
|
||||
|
||||
it('should support named access for HFields', () => {
|
||||
const xml = `<mujoco model="test_hfield">
|
||||
<asset>
|
||||
<hfield name="hf" nrow="2" ncol="3" size="1 2 0.1 0.2"/>
|
||||
</asset>
|
||||
<worldbody>
|
||||
<geom type="hfield" hfield="hf"/>
|
||||
</worldbody>
|
||||
</mujoco>`;
|
||||
const tempXmlFilename = '/tmp/hf.xml';
|
||||
writeXMLFile(tempXmlFilename, xml);
|
||||
const model = mujoco.MjModel.mj_loadXML(tempXmlFilename);
|
||||
try {
|
||||
assertExists(model);
|
||||
const hf = model.hfield('hf');
|
||||
expect(hf.id).toBe(0);
|
||||
expect(hf.name).toBe('hf');
|
||||
expect(hf.nrow).toBe(2);
|
||||
expect(hf.ncol).toBe(3);
|
||||
expectArraysClose(hf.size, new Float64Array([1, 2, 0.1, 0.2]));
|
||||
const expectedDataLength =
|
||||
model.hfield_nrow[hf.id] * model.hfield_ncol[hf.id];
|
||||
expect(expectedDataLength).toBe(6);
|
||||
expect(hf.data.length).toBe(expectedDataLength);
|
||||
} finally {
|
||||
model?.delete();
|
||||
unlinkXMLFile(tempXmlFilename);
|
||||
}
|
||||
});
|
||||
|
||||
it('should support named access for Numeric', () => {
|
||||
const xml = `<mujoco>
|
||||
<custom>
|
||||
<numeric name="n1" data="1 2 3"/>
|
||||
<numeric name="n2" data="4 5 6 7"/>
|
||||
</custom>
|
||||
</mujoco>`;
|
||||
const tempXmlFilename = '/tmp/num.xml';
|
||||
writeXMLFile(tempXmlFilename, xml);
|
||||
const model = mujoco.MjModel.mj_loadXML(tempXmlFilename);
|
||||
try {
|
||||
assertExists(model);
|
||||
const n1 = model.numeric('n1');
|
||||
const expectedN1data = model.numeric_data.slice(
|
||||
model.numeric_adr[n1.id],
|
||||
model.numeric_adr[n1.id] + model.numeric_size[n1.id]);
|
||||
expect(n1.id).toBe(0);
|
||||
expect(n1.name).toBe('n1');
|
||||
expect(n1.size).toBe(3);
|
||||
expect(expectedN1data).toEqual(new Float64Array([1, 2, 3]));
|
||||
expectArraysClose(n1.data, expectedN1data);
|
||||
|
||||
const n2 = model.numeric('n2');
|
||||
const expectedN2data = model.numeric_data.slice(
|
||||
model.numeric_adr[n2.id],
|
||||
model.numeric_adr[n2.id] + model.numeric_size[n2.id]);
|
||||
expect(n2.id).toBe(1);
|
||||
expect(n2.name).toBe('n2');
|
||||
expect(n2.size).toBe(4);
|
||||
expect(expectedN2data).toEqual(new Float64Array([4, 5, 6, 7]));
|
||||
expectArraysClose(n2.data, expectedN2data);
|
||||
} finally {
|
||||
model?.delete();
|
||||
unlinkXMLFile(tempXmlFilename);
|
||||
}
|
||||
});
|
||||
|
||||
it('should support named access for Texture', () => {
|
||||
const xml = `<mujoco>
|
||||
<asset>
|
||||
<texture name="t1" type="2d" builtin="checker" width="10"
|
||||
height="12"/>
|
||||
</asset>
|
||||
</mujoco>`;
|
||||
const tempXmlFilename = '/tmp/tex.xml';
|
||||
writeXMLFile(tempXmlFilename, xml);
|
||||
const model = mujoco.MjModel.mj_loadXML(tempXmlFilename);
|
||||
try {
|
||||
assertExists(model);
|
||||
const t1 = model.tex('t1');
|
||||
const expectedData = model.tex_data.slice(
|
||||
model.tex_adr[t1.id],
|
||||
model.tex_adr[t1.id] +
|
||||
model.tex_height[t1.id] * model.tex_width[t1.id] *
|
||||
model.tex_nchannel[t1.id]);
|
||||
|
||||
expect(t1.id).toBe(0);
|
||||
expect(t1.name).toBe('t1');
|
||||
expect(t1.type).toBe(mujoco.mjtTexture.mjTEXTURE_2D.value);
|
||||
expect(t1.width).toBe(10);
|
||||
expect(t1.height).toBe(12);
|
||||
expect(t1.nchannel).toBe(3);
|
||||
expect(t1.data.length).toBe(t1.width * t1.height * t1.nchannel);
|
||||
expectArraysClose(t1.data, expectedData);
|
||||
} finally {
|
||||
model?.delete();
|
||||
unlinkXMLFile(tempXmlFilename);
|
||||
}
|
||||
});
|
||||
|
||||
it('should support named access for Tuple', () => {
|
||||
const xml = `<mujoco>
|
||||
<worldbody>
|
||||
<geom name="g1" type="sphere" size=".1"/>
|
||||
<geom name="g2" type="box" size=".1 .1 .1"/>
|
||||
</worldbody>
|
||||
<custom>
|
||||
<tuple name="tup1">
|
||||
<element objtype="geom" objname="g1" prm="0.1"/>
|
||||
<element objtype="geom" objname="g2" prm="0.2"/>
|
||||
</tuple>
|
||||
</custom>
|
||||
</mujoco>`;
|
||||
const tempXmlFilename = '/tmp/tup.xml';
|
||||
writeXMLFile(tempXmlFilename, xml);
|
||||
const model = mujoco.MjModel.mj_loadXML(tempXmlFilename);
|
||||
try {
|
||||
assertExists(model);
|
||||
const t1 = model.tuple('tup1');
|
||||
const expectedObjtype = model.tuple_objtype.slice(
|
||||
model.tuple_adr[t1.id],
|
||||
model.tuple_adr[t1.id] + model.tuple_size[t1.id]);
|
||||
expect(t1.id).toBe(0);
|
||||
expect(t1.name).toBe('tup1');
|
||||
expect(t1.size).toBe(2);
|
||||
expectArraysEqual(
|
||||
expectedObjtype, new Int32Array([
|
||||
mujoco.mjtObj.mjOBJ_GEOM.value, mujoco.mjtObj.mjOBJ_GEOM.value
|
||||
]));
|
||||
expectArraysEqual(t1.objtype, expectedObjtype);
|
||||
} finally {
|
||||
model?.delete();
|
||||
unlinkXMLFile(tempXmlFilename);
|
||||
}
|
||||
});
|
||||
});
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user