Refactor MuJoCo WASM bindings functions code generation and auto-generate functions requiring argument boundchecks.

PiperOrigin-RevId: 832333773
Change-Id: Id2e1a59e8d366baa2ee044d693f76aadcd1dfd32
This commit is contained in:
Matias Manevi
2025-11-14 08:50:04 -08:00
committed by Copybara-Service
parent 3f69deb328
commit a36f452250
5 changed files with 1256 additions and 1791 deletions
File diff suppressed because it is too large Load Diff
+327 -86
View File
@@ -187,6 +187,16 @@ _SKIPPED_UTILITY_FUNCTIONS: List[str] = [
# go/keep-sorted end
]
# Functions that require special wrappers.
# These functions are not bound automatically but are written by hand instead.
MANUAL_WRAPPER_FUNCTIONS: List[str] = [
# go/keep-sorted start
"mj_saveLastXML",
"mj_setLengthRange",
"mju_error",
# go/keep-sorted end
]
# List of functions that should be skipped during the code generation process.
SKIPPED_FUNCTIONS: List[str] = (
_SKIPPED_CLASS_METHODS
@@ -201,92 +211,6 @@ SKIPPED_FUNCTIONS: List[str] = (
+ _SKIPPED_UTILITY_FUNCTIONS
)
# Functions that require special wrappers to infer sizes and make additional
# validation checks. These functions are not bound automatically but are
# written by hand instead.
BOUNDCHECK_FUNCS: List[str] = [
# go/keep-sorted start
"mj_addM",
"mj_angmomMat",
"mj_applyFT",
"mj_constraintUpdate",
"mj_differentiatePos",
"mj_fullM",
"mj_geomDistance",
"mj_getState",
"mj_integratePos",
"mj_jac",
"mj_jacBody",
"mj_jacBodyCom",
"mj_jacDot",
"mj_jacGeom",
"mj_jacPointAxis",
"mj_jacSite",
"mj_jacSubtreeCom",
"mj_mulJacTVec",
"mj_mulJacVec",
"mj_mulM",
"mj_mulM2",
"mj_multiRay",
"mj_normalizeQuat",
"mj_rne",
"mj_saveLastXML",
"mj_setLengthRange",
"mj_setState",
"mj_solveM",
"mj_solveM2",
"mjd_inverseFD",
"mjd_subQuat",
"mjd_transitionFD",
"mju_L1",
"mju_add",
"mju_addScl",
"mju_addTo",
"mju_addToScl",
"mju_band2Dense",
"mju_bandMulMatVec",
"mju_boxQP",
"mju_cholFactor",
"mju_cholFactorBand",
"mju_cholSolve",
"mju_cholSolveBand",
"mju_cholUpdate",
"mju_copy",
"mju_d2n",
"mju_decodePyramid",
"mju_dense2Band",
"mju_dense2sparse",
"mju_dot",
"mju_encodePyramid",
"mju_eye",
"mju_f2n",
"mju_fill",
"mju_insertionSort",
"mju_insertionSortInt",
"mju_isZero",
"mju_mulMatMat",
"mju_mulMatMatT",
"mju_mulMatTMat",
"mju_mulMatTVec",
"mju_mulMatVec",
"mju_mulVecMatVec",
"mju_n2d",
"mju_n2f",
"mju_norm",
"mju_normalize",
"mju_printMatSparse",
"mju_scl",
"mju_sparse2dense",
"mju_sqrMatTD",
"mju_sub",
"mju_subFrom",
"mju_sum",
"mju_symmetrize",
"mju_transpose",
"mju_zero",
# go/keep-sorted end
]
# List of structs that should be skipped during the code generation process.
SKIPPED_STRUCTS: List[str] = [
# go/keep-sorted start
@@ -480,3 +404,320 @@ BYTE_FIELDS: Dict[str, Dict[str, str]] = {
"buffer": {"size": "nbuffer"},
"arena": {"size": "narena"},
}
# pyformat: disable
# Dictionary mapping function names to their boundcheck code.
FUNCTION_BOUNDS_CHECKS: Dict[str, str] = {
"mj_solveM": """
CHECK_SIZES(x, y);
CHECK_DIVISIBLE(x, m.nv());
int n = x_div.quot;
""".strip(),
"mj_solveM2": """
CHECK_SIZES(x, y);
CHECK_SIZE(sqrtInvD, m.nv());
CHECK_DIVISIBLE(x, m.nv());
int n = x_div.quot;
""".strip(),
"mju_add": """
CHECK_SIZES(res, vec1);
CHECK_SIZES(res, vec2);
int n = res_.size();
""".strip(),
"mju_addScl": """
CHECK_SIZES(res, vec1);
CHECK_SIZES(res, vec2);
int n = res_.size();
""".strip(),
"mju_addTo": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_addToScl": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_boxQP": """
CHECK_SIZES(lower, res);
CHECK_SIZES(upper, res);
CHECK_SIZES(index, res);
CHECK_SIZE(R, res_.size() * (res_.size() + 7))
CHECK_PERFECT_SQUARE(H);
CHECK_SIZES(g, res);
int n = res_.size();
""".strip(),
"mju_cholFactor": """
CHECK_PERFECT_SQUARE(mat);
int n = mat_sqrt;
""".strip(),
"mju_cholSolve": """
CHECK_PERFECT_SQUARE(mat);
CHECK_SIZE(res, mat_sqrt);
CHECK_SIZE(vec, mat_sqrt);
int n = mat_sqrt;
""".strip(),
"mju_cholUpdate": """
CHECK_PERFECT_SQUARE(mat);
CHECK_SIZE(x, mat_sqrt);
int n = mat_sqrt;
""".strip(),
"mju_copy": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_d2n": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_decodePyramid": """
CHECK_SIZE(pyramid, 2 * mu_.size());
CHECK_SIZE(force, mu_.size() + 1);
int dim = mu_.size();
""".strip(),
"mju_dot": """
CHECK_SIZES(vec1, vec2);
int n = vec1_.size();
""".strip(),
"mju_encodePyramid": """
CHECK_SIZE(pyramid, 2 * mu_.size());
CHECK_SIZE(force, mu_.size() + 1);
int dim = mu_.size();
""".strip(),
"mju_eye": """
CHECK_PERFECT_SQUARE(mat);
int n = mat_sqrt;
""".strip(),
"mju_f2n": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_mulVecMatVec": """
CHECK_SIZES(vec1, vec2);
CHECK_SIZE(mat, vec1_.size() * vec2_.size());
int n = vec1_.size();
""".strip(),
"mju_n2d": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_n2f": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_printMatSparse": """
CHECK_SIZES(rownnz, rowadr);
int nr = rowadr_.size();
""".strip(),
"mju_scl": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_sub": """
CHECK_SIZES(res, vec1);
CHECK_SIZES(res, vec2);
int n = res_.size();
""".strip(),
"mju_subFrom": """
CHECK_SIZES(res, vec);
int n = res_.size();
""".strip(),
"mju_insertionSort": "int n = list_.size();",
"mju_insertionSortInt": "int n = list_.size();",
"mju_fill": "int n = res_.size();",
"mju_dense2sparse": """
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(rownnz, nr);
CHECK_SIZE(rowadr, nr);
CHECK_SIZE(colind, res_.size());
int nnz = res_.size();
""".strip(),
"mj_addM": """
CHECK_SIZE(rownnz, m.nv());
CHECK_SIZE(rowadr, m.nv());
CHECK_SIZE(colind, m.nM());
CHECK_SIZE(dst, m.nM());
""".strip(),
"mj_angmomMat": """
CHECK_SIZE(mat, m.nv() * 3);
""".strip(),
"mj_applyFT": """
CHECK_SIZE(qfrc_target, m.nv());
CHECK_SIZE(force, 3);
CHECK_SIZE(torque, 3);
CHECK_SIZE(point, 3);
""".strip(),
"mj_constraintUpdate": """
CHECK_SIZE(cost, 1);
CHECK_SIZE(jar, d.nefc());
""".strip(),
"mj_differentiatePos": """
CHECK_SIZE(qvel, m.nv());
CHECK_SIZE(qpos1, m.nq());
CHECK_SIZE(qpos2, m.nq());
""".strip(),
"mj_fullM": """
CHECK_SIZE(M, m.nM());
CHECK_SIZE(dst, m.nv() * m.nv());
""".strip(),
"mj_geomDistance": """
CHECK_SIZE(fromto, 6);
""".strip(),
"mj_getState": """
CHECK_SIZE(state, mj_stateSize(m.get(), sig));
""".strip(),
"mj_integratePos": """
CHECK_SIZE(qpos, m.nq());
CHECK_SIZE(qvel, m.nv());
""".strip(),
"mj_jac": """
CHECK_SIZE(point, 3);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
""".strip(),
"mj_jacBody": """
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
""".strip(),
"mj_jacBodyCom": """
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
""".strip(),
"mj_jacDot": """
CHECK_SIZE(point, 3);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
""".strip(),
"mj_jacGeom": """
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
""".strip(),
"mj_jacPointAxis": """
CHECK_SIZE(point, 3);
CHECK_SIZE(axis, 3);
CHECK_SIZE(jacPoint, m.nv() * 3);
CHECK_SIZE(jacAxis, m.nv() * 3);
""".strip(),
"mj_jacSite": """
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
""".strip(),
"mj_jacSubtreeCom": """
CHECK_SIZE(jacp, m.nv() * 3);
""".strip(),
"mj_mulJacTVec": """
CHECK_SIZE(res, m.nv());
CHECK_SIZE(vec, d.nefc());
""".strip(),
"mj_mulJacVec": """
CHECK_SIZE(res, d.nefc());
CHECK_SIZE(vec, m.nv());
""".strip(),
"mj_mulM": """
CHECK_SIZE(res, m.nv());
CHECK_SIZE(vec, m.nv());
""".strip(),
"mj_mulM2": """
CHECK_SIZE(res, m.nv());
CHECK_SIZE(vec, m.nv());
""".strip(),
"mj_multiRay": """
CHECK_SIZE(dist, nray);
CHECK_SIZE(geomid, nray);
CHECK_SIZE(vec, 3 * nray);
""".strip(),
"mj_normalizeQuat": """
CHECK_SIZE(qpos, m.nq());
""".strip(),
"mj_rne": """
CHECK_SIZE(result, m.nv());
""".strip(),
"mj_setState": """
CHECK_SIZE(state, mj_stateSize(m.get(), sig));
""".strip(),
"mjd_inverseFD": """
CHECK_SIZE(DfDq, m.nv() * m.nv());
CHECK_SIZE(DfDv, m.nv() * m.nv());
CHECK_SIZE(DfDa, m.nv() * m.nv());
CHECK_SIZE(DsDq, m.nv() * m.nsensordata());
CHECK_SIZE(DsDv, m.nv() * m.nsensordata());
CHECK_SIZE(DsDa, m.nv() * m.nsensordata());
CHECK_SIZE(DmDq, m.nv() * m.nM());
""".strip(),
"mjd_subQuat": """
CHECK_SIZE(qa, 4);
CHECK_SIZE(qb, 4);
CHECK_SIZE(Da, 9);
CHECK_SIZE(Db, 9);
""".strip(),
"mjd_transitionFD": """
CHECK_SIZE(A, (2 * m.nv() + m.na()) * (2 * m.nv() + m.na()));
CHECK_SIZE(B, (2 * m.nv() + m.na()) * m.nu());
CHECK_SIZE(C, m.nsensordata() * (2 * m.nv() + m.na()));
CHECK_SIZE(D, m.nsensordata() * m.nu());
""".strip(),
"mju_band2Dense": """
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
CHECK_SIZE(res, ntotal * ntotal);
""".strip(),
"mju_bandMulMatVec": """
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
CHECK_SIZE(res, ntotal * nvec);
CHECK_SIZE(vec, ntotal * nvec);
""".strip(),
"mju_cholFactorBand": """
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
""".strip(),
"mju_cholSolveBand": """
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
CHECK_SIZE(res, ntotal);
CHECK_SIZE(vec, ntotal);
""".strip(),
"mju_dense2Band": """
CHECK_SIZE(mat, ntotal * ntotal);
CHECK_SIZE(res, (ntotal - ndense) * nband + ndense * ntotal);
""".strip(),
"mju_mulMatMat": """
CHECK_SIZE(res, r1 * c2);
CHECK_SIZE(mat1, r1 * c1);
CHECK_SIZE(mat2, c1 * c2);
""".strip(),
"mju_mulMatMatT": """
CHECK_SIZE(res, r1 * r2);
CHECK_SIZE(mat1, r1 * c1);
CHECK_SIZE(mat2, r2 * c1);
""".strip(),
"mju_mulMatTMat": """
CHECK_SIZE(res, c1 * c2);
CHECK_SIZE(mat1, r1 * c1);
CHECK_SIZE(mat2, r1 * c2);
""".strip(),
"mju_mulMatTVec": """
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nc);
CHECK_SIZE(vec, nr);
""".strip(),
"mju_mulMatVec": """
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nr);
CHECK_SIZE(vec, nc);
""".strip(),
"mju_sparse2dense": """
CHECK_SIZE(res, nr * nc);
CHECK_SIZE(rownnz, nr);
CHECK_SIZE(rowadr, nr);
""".strip(),
"mju_sqrMatTD": """
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nc * nc);
CHECK_SIZE(diag, nr);
""".strip(),
"mju_symmetrize": """
CHECK_SIZE(mat, n * n);
CHECK_SIZE(res, n * n);
""".strip(),
"mju_transpose": """
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nr * nc);
""".strip(),
}
# pyformat: enable
+104 -102
View File
@@ -114,129 +114,131 @@ def should_be_wrapped(func: ast_nodes.FunctionDecl) -> bool:
def generate_function_wrapper(func: ast_nodes.FunctionDecl) -> str:
"""Generates C++ code for a wrapper function."""
params_unpack_statements = get_params_unpack_statements(func.parameters)
wrapper_params_list = get_params_string(func.parameters)
not_nullable_params = get_params_notnullable(func.parameters)
wrapper_params = ", ".join(wrapper_params_list)
bound_check_code = constants.FUNCTION_BOUNDS_CHECKS.get(func.name, "")
wrapper_parameters = [
p for p in func.parameters if f"{p.name} = " not in bound_check_code
]
wrapper_params = ", ".join([get_param_string(p) for p in wrapper_parameters])
ret_type = get_compatible_return_type(func)
builder = code_builder.CodeBuilder()
with builder.function(f"{ret_type} {func.name}_wrapper({wrapper_params})"):
invoker_params_list = get_params_string_maybe_with_conversion(
for p in wrapper_parameters:
if c_notnullable := get_param_notnullable(p):
builder.line(f"CHECK_VAL({c_notnullable});")
for p in wrapper_parameters:
if c_unpack := get_param_unpack_statement(p):
builder.line(c_unpack)
if bound_check_code:
builder.line(bound_check_code)
c_params_list = get_params_string_maybe_with_conversion(
func.parameters
)
invoker_params_str = ", ".join(invoker_params_list)
invoker_call = f"{func.name}({invoker_params_str})"
invoker_statement = get_compatible_return_call(func, invoker_call)
for p in not_nullable_params:
builder.line(f"CHECK_VAL({p});")
for unpack_statement in params_unpack_statements:
builder.line(unpack_statement)
builder.line(f"{invoker_statement};")
c_params_str = ", ".join(c_params_list)
c_call = f"{func.name}({c_params_str})"
c_statement = get_compatible_return_call(func, c_call)
builder.line(f"{c_statement};")
return builder.to_string()
def get_params_notnullable(
ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...],
) -> List[str]:
def get_param_notnullable(
p: ast_nodes.FunctionParameterDecl,
) -> str:
"""Generates list of param names for checking if they aren't null/undefined."""
not_nullable_params = []
for p in ast_params:
if (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
# We only check for char because others are checked in the unpacker
# and we don't want to check twice.
and p.type.inner_type.name == "char"
and not p.nullable
):
not_nullable_params.append(p.name)
return not_nullable_params
if (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
# We only check for char because others are checked in the unpacker
# and we don't want to check twice.
and p.type.inner_type.name == "char"
and not p.nullable
):
return p.name
return ""
def get_params_unpack_statements(
ast_params: Tuple[ast_nodes.FunctionParameterDecl, ...],
) -> List[str]:
def get_param_unpack_statement(
p: ast_nodes.FunctionParameterDecl,
) -> str | None:
"""Generates C++ statements to unpack JS values for pointer/array parameters."""
params_unpack_statements = []
for p in ast_params:
if (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
):
if p.type.inner_type.name == "char":
# param is Javascript string
continue
if p.type.inner_type.is_const:
# param is Javascript number[]
params_unpack_statements.append(
f"UNPACK_ARRAY({p.type.inner_type.name}, {p.name});"
)
if (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
):
if p.type.inner_type.name == "char":
# param is Javascript string
return ""
if p.type.inner_type.is_const:
# param is Javascript number[]
if p.nullable:
return f"UNPACK_NULLABLE_ARRAY({p.type.inner_type.name}, {p.name});"
else:
# param is TypedArray or a WasmBuffer
params_unpack_statements.append(
f"UNPACK_VALUE({p.type.inner_type.name}, {p.name});"
)
return params_unpack_statements
return f"UNPACK_ARRAY({p.type.inner_type.name}, {p.name});"
else:
# param is TypedArray or a WasmBuffer
if p.nullable:
return f"UNPACK_NULLABLE_VALUE({p.type.inner_type.name}, {p.name});"
else:
return f"UNPACK_VALUE({p.type.inner_type.name}, {p.name});"
return None
def get_params_string(
parameters: Tuple[ast_nodes.FunctionParameterDecl, ...]
) -> List[str]:
def get_param_string(
p: ast_nodes.FunctionParameterDecl
) -> str:
"""Generates a list of C++ parameter declarations as strings."""
result = []
for p in parameters:
if (
isinstance(p.type, ast_nodes.PointerType)
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name not in PRIMITIVE_TYPES
):
# Pointer to struct parameters
const_qualifier = "const " if p.type.inner_type.is_const else ""
result.append(
f"{const_qualifier}{uppercase_first_letter(p.type.inner_type.name)}&"
f" {p.name}"
)
elif (
isinstance(p.type, ast_nodes.ValueType)
and p.type.name in PRIMITIVE_TYPES
):
# Primitive value parameters
const_qualifier = "const " if p.type.is_const else ""
result.append(f"{const_qualifier}{p.type} {p.name}")
elif (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
):
# Pointer to primitive value parameters or arrays
if p.type.inner_type.name == "char":
if p.nullable:
result.append(f"const NullableString& {p.name}")
else:
result.append(f"const String& {p.name}")
elif (
p.type.inner_type.name
in ["int", "float", "double", "mjtNum", "mjtByte"]
and p.type.inner_type.is_const
):
result.append(f"const NumberArray& {p.name}")
if (
isinstance(p.type, ast_nodes.PointerType)
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name not in PRIMITIVE_TYPES
):
# Pointer to struct parameters
const_qualifier = "const " if p.type.inner_type.is_const else ""
return (
f"{const_qualifier}{uppercase_first_letter(p.type.inner_type.name)}&"
f" {p.name}"
)
elif (
isinstance(p.type, ast_nodes.ValueType)
and p.type.name in PRIMITIVE_TYPES
):
# Primitive value parameters
const_qualifier = "const " if p.type.is_const else ""
return f"{const_qualifier}{p.type} {p.name}"
elif (
isinstance(p.type, (ast_nodes.PointerType, ast_nodes.ArrayType))
and isinstance(p.type.inner_type, ast_nodes.ValueType)
and p.type.inner_type.name in PRIMITIVE_TYPES
):
# Pointer to primitive value parameters or arrays
if p.type.inner_type.name == "char":
if p.nullable:
return f"const NullableString& {p.name}"
else:
result.append(f"const val& {p.name}")
return f"const String& {p.name}"
elif (
p.type.inner_type.name
in ["int", "float", "double", "mjtNum", "mjtByte"]
and p.type.inner_type.is_const
):
return f"const NumberArray& {p.name}"
else:
# This case should ideally not be reached if AST is well-formed
# and types are categorized by the helper booleans correctly.
raise TypeError(
"Unable to generate param string. Unhandled parameter type:"
f" {p.type} for param '{p.name}'"
)
return result
return f"const val& {p.name}"
else:
# This case should ideally not be reached if AST is well-formed
# and types are categorized by the helper booleans correctly.
raise TypeError(
"Unable to generate param string. Unhandled parameter type:"
f" {p.type} for param '{p.name}'"
)
def get_params_string_maybe_with_conversion(
@@ -366,7 +368,7 @@ class Generator:
code = []
for func in self.wrapper_bind_functions:
if func.name not in constants.BOUNDCHECK_FUNCS:
if func.name not in constants.MANUAL_WRAPPER_FUNCTIONS:
wrapper_code = generate_function_wrapper(func)
code.append(wrapper_code)
+68 -789
View File
@@ -38,6 +38,73 @@
namespace mujoco::wasm {
using emscripten::enum_;
using emscripten::class_;
using emscripten::function;
using emscripten::val;
using emscripten::constant;
using emscripten::register_optional;
using emscripten::register_type;
using emscripten::register_vector;
using emscripten::return_value_policy::reference;
using emscripten::return_value_policy::take_ownership;
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray);
EMSCRIPTEN_DECLARE_VAL_TYPE(String);
// 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?
#define CHECK_VAL(val) \
if (val.isNull()) { \
mju_error("Invalid argument: %s is null", #val); \
} else if (val.isUndefined()) { \
mju_error("Invalid argument: %s is undefined", #val); \
}
void ThrowMujocoErrorToJS(const char* msg) {
// Get a handle to the JS global Error constructor function, create a new
// object instance and then throw the object as an exception using the
// val::throw_() helper function.
val(val::global("Error").new_(val("MuJoCo Error: " + std::string(msg))))
.throw_();
}
__attribute__((constructor)) void InitMuJoCoErrorHandler() {
mju_user_error = ThrowMujocoErrorToJS;
}
template <size_t N>
val MakeValArray(const char* (&strings)[N]) {
val result = val::array();
for (int i = 0; i < N; i++) {
result.call<void>("push", val(strings[i]));
}
return result;
}
template <size_t N, size_t M>
val MakeValArray3(const char* (&strings)[N][M]) {
val result = val::array();
for (int i = 0; i < N; i++) {
val inner = val::array();
for (int j = 0; j < M; j++) {
inner.call<void>("push", val(strings[i][j]));
}
result.call<void>("push", inner);
}
return result;
}
template <typename WrapperType, typename ArrayType, typename SizeType>
std::vector<WrapperType> InitWrapperArray(ArrayType* array, SizeType size) {
std::vector<WrapperType> result;
result.reserve(size);
for (int i = 0; i < size; ++i) {
result.emplace_back(&array[i]);
}
return result;
}
// Create the types for anonymous structs
using mjVisualGlobal = decltype(::mjVisual::global);
using mjVisualQuality = decltype(::mjVisual::quality);
@@ -230,52 +297,6 @@ struct MjSpec {
MjsElement element;
};
using emscripten::enum_;
using emscripten::class_;
using emscripten::function;
using emscripten::val;
using emscripten::constant;
using emscripten::register_optional;
using emscripten::register_type;
using emscripten::register_vector;
using emscripten::return_value_policy::reference;
using emscripten::return_value_policy::take_ownership;
// ERROR HANDLER
void ThrowMujocoErrorToJS(const char* msg) {
// Get a handle to the JS global Error constructor function, create a new
// object instance and then throw the object as an exception using the
// val::throw_() helper function.
val(val::global("Error").new_(val("MuJoCo Error: " + std::string(msg))))
.throw_();
}
__attribute__((constructor)) void InitMuJoCoErrorHandler() {
mju_user_error = ThrowMujocoErrorToJS;
}
// CONSTANTS
template <size_t N>
val MakeValArray(const char* (&strings)[N]) {
val result = val::array();
for (int i = 0; i < N; i++) {
result.call<void>("push", val(strings[i]));
}
return result;
}
template <size_t N, size_t M>
val MakeValArray3(const char* (&strings)[N][M]) {
val result = val::array();
for (int i = 0; i < N; i++) {
val inner = val::array();
for (int j = 0; j < M; j++) {
inner.call<void>("push", val(strings[i][j]));
}
result.call<void>("push", inner);
}
return result;
}
val get_mjDISABLESTRING() { return MakeValArray(mjDISABLESTRING); }
val get_mjENABLESTRING() { return MakeValArray(mjENABLESTRING); }
val get_mjTIMERSTRING() { return MakeValArray(mjTIMERSTRING); }
@@ -330,17 +351,7 @@ EMSCRIPTEN_BINDINGS(mujoco_enums) {
// {{ ENUM_BINDINGS }}
}
// STRUCTS
// {{ AUTOGENNED_STRUCTS_SOURCE }}
template <typename WrapperType, typename ArrayType, typename SizeType>
std::vector<WrapperType> InitWrapperArray(ArrayType* array, SizeType size) {
std::vector<WrapperType> result;
result.reserve(size);
for (int i = 0; i < size; ++i) {
result.emplace_back(&array[i]);
}
return result;
}
// =============== MjModel =============== //
MjModel::MjModel(mjModel *m)
@@ -544,64 +555,9 @@ EMSCRIPTEN_BINDINGS(mujoco_structs) {
emscripten::register_vector<MjvGeom>("MjvGeomVec");
}
// FUNCTIONS
EMSCRIPTEN_DECLARE_VAL_TYPE(NumberArray);
EMSCRIPTEN_DECLARE_VAL_TYPE(String);
// 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?
#define CHECK_VAL(val) \
if (val.isNull()) { \
mju_error("Invalid argument: %s is null", #val); \
} else if (val.isUndefined()) { \
mju_error("Invalid argument: %s is undefined", #val); \
}
void error_wrapper(const String& msg) { mju_error("%s\n", msg.as<const std::string>().data()); }
// {{ WRAPPER_FUNCTIONS }}
void mju_printMatSparse_wrapper(const NumberArray& mat, const NumberArray& rownnz, const NumberArray& rowadr, const NumberArray& colind)
{
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(int, rownnz);
UNPACK_ARRAY(int, rowadr);
UNPACK_ARRAY(int, colind);
CHECK_SIZES(rownnz, rowadr);
mju_printMatSparse(mat_.data(), rowadr_.size(),
rownnz_.data(),
rowadr_.data(),
colind_.data());
}
void mj_solveM_wrapper(const MjModel& m, MjData& d, const val& x, const NumberArray& y)
{
UNPACK_VALUE(mjtNum, x);
UNPACK_ARRAY(mjtNum, y);
CHECK_SIZES(x, y);
CHECK_DIVISIBLE(x, m.nv());
mj_solveM(m.get(), d.get(), x_.data(), y_.data(), x_div.quot);
}
void mj_solveM2_wrapper(const MjModel& m, MjData& d,
const val& x, const NumberArray& y,
const NumberArray& sqrtInvD) {
UNPACK_VALUE(mjtNum, x);
UNPACK_ARRAY(mjtNum, y);
UNPACK_ARRAY(mjtNum, sqrtInvD);
CHECK_SIZES(x, y);
CHECK_SIZE(sqrtInvD, m.nv());
CHECK_DIVISIBLE(x, m.nv());
mj_solveM2(m.get(), d.get(), x_.data(), y_.data(), sqrtInvD_.data(), x_div.quot);
}
void mj_rne_wrapper(const MjModel& m, MjData& d, int flg_acc, const val& result)
{
UNPACK_VALUE(mjtNum, result);
CHECK_SIZE(result, m.nv());
mj_rne(m.get(), d.get(), flg_acc, result_.data());
}
void error_wrapper(const String& msg) { mju_error("%s\n", msg.as<const std::string>().data()); }
int mj_saveLastXML_wrapper(const String& filename, const MjModel& m) {
CHECK_VAL(filename);
@@ -622,683 +578,6 @@ int mj_setLengthRange_wrapper(const MjModel& m, const MjData& d, int index, cons
return result;
}
void mj_constraintUpdate_wrapper(const MjModel& m, MjData& d, const NumberArray& jar, const val& cost, int flg_coneHessian)
{
UNPACK_ARRAY(mjtNum, jar);
UNPACK_NULLABLE_VALUE(mjtNum, cost);
CHECK_SIZE(cost, 1);
CHECK_SIZE(jar, d.nefc());
mj_constraintUpdate(m.get(), d.get(), jar_.data(), cost_.data(), flg_coneHessian);
}
void mj_getState_wrapper(const MjModel& m, const MjData& d, const val& state, unsigned int spec)
{
UNPACK_VALUE(mjtNum, state);
CHECK_SIZE(state, mj_stateSize(m.get(), spec));
mj_getState(m.get(), d.get(), state_.data(), spec);
}
void mj_setState_wrapper(const MjModel& m, MjData& d, const NumberArray& state, unsigned int spec)
{
UNPACK_ARRAY(mjtNum, state);
CHECK_SIZE(state, mj_stateSize(m.get(), spec));
mj_setState(m.get(), d.get(), state_.data(), spec);
}
void mj_mulJacVec_wrapper(const MjModel& m, const MjData& d, const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(res, d.nefc());
CHECK_SIZE(vec, m.nv());
mj_mulJacVec(m.get(), d.get(), res_.data(), vec_.data());
}
void mj_mulJacTVec_wrapper(const MjModel& m, const MjData& d, const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(res, m.nv());
CHECK_SIZE(vec, d.nefc());
mj_mulJacTVec(m.get(), d.get(), res_.data(), vec_.data());
}
void mj_jac_wrapper(const MjModel& m, const MjData& d, const val& jacp, const val& jacr, const NumberArray& point, int body)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacp);
UNPACK_NULLABLE_VALUE(mjtNum, jacr);
UNPACK_ARRAY(mjtNum, point);
CHECK_SIZE(point, 3);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
mj_jac(m.get(), d.get(), jacp_.data(), jacr_.data(), point_.data(), body);
}
void mj_jacBody_wrapper(const MjModel& m, const MjData& d, const val& jacp, const val& jacr, int body)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacp);
UNPACK_NULLABLE_VALUE(mjtNum, jacr);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
mj_jacBody(m.get(), d.get(), jacp_.data(), jacr_.data(), body);
}
void mj_jacBodyCom_wrapper(const MjModel& m, const MjData& d, const val& jacp, const val& jacr, int body)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacp);
UNPACK_NULLABLE_VALUE(mjtNum, jacr);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
mj_jacBodyCom(m.get(), d.get(), jacp_.data(), jacr_.data(), body);
}
void mj_jacSubtreeCom_wrapper(const MjModel& m, MjData& d, const val& jacp, int body)
{
UNPACK_VALUE(mjtNum, jacp);
CHECK_SIZE(jacp, m.nv() * 3);
mj_jacSubtreeCom(m.get(), d.get(), jacp_.data(), body);
}
void mj_jacGeom_wrapper(const MjModel& m, const MjData& d, const val& jacp, const val& jacr, int geom)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacp);
UNPACK_NULLABLE_VALUE(mjtNum, jacr);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
mj_jacGeom(m.get(), d.get(), jacp_.data(), jacr_.data(), geom);
}
void mj_jacSite_wrapper(const MjModel& m, const MjData& d, const val& jacp, const val& jacr, int site)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacp);
UNPACK_NULLABLE_VALUE(mjtNum, jacr);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
mj_jacSite(m.get(), d.get(), jacp_.data(), jacr_.data(), site);
}
void mj_jacPointAxis_wrapper(const MjModel& m, MjData& d, const val& jacPoint, const val& jacAxis, const NumberArray& point, const NumberArray& axis, int body)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacPoint);
UNPACK_NULLABLE_VALUE(mjtNum, jacAxis);
UNPACK_ARRAY(mjtNum, point);
UNPACK_ARRAY(mjtNum, axis);
CHECK_SIZE(point, 3);
CHECK_SIZE(axis, 3);
CHECK_SIZE(jacPoint, m.nv() * 3);
CHECK_SIZE(jacAxis, m.nv() * 3);
mj_jacPointAxis(m.get(), d.get(), jacPoint_.data(), jacAxis_.data(), point_.data(), axis_.data(), body);
}
void mj_jacDot_wrapper(const MjModel& m, const MjData& d, const val& jacp, const val& jacr, const NumberArray& point, int body)
{
UNPACK_NULLABLE_VALUE(mjtNum, jacp);
UNPACK_NULLABLE_VALUE(mjtNum, jacr);
UNPACK_ARRAY(mjtNum, point);
CHECK_SIZE(point, 3);
CHECK_SIZE(jacp, m.nv() * 3);
CHECK_SIZE(jacr, m.nv() * 3);
mj_jacDot(m.get(), d.get(), jacp_.data(), jacr_.data(), point_.data(), body);
}
void mj_angmomMat_wrapper(const MjModel& m, MjData& d, const val& mat, int body)
{
UNPACK_VALUE(mjtNum, mat);
CHECK_SIZE(mat, m.nv() * 3);
mj_angmomMat(m.get(), d.get(), mat_.data(), body);
}
void mj_fullM_wrapper(const MjModel& m, const val& dst, const NumberArray& M)
{
UNPACK_VALUE(mjtNum, dst);
UNPACK_ARRAY(mjtNum, M);
CHECK_SIZE(M, m.nM());
CHECK_SIZE(dst, m.nv() * m.nv());
mj_fullM(m.get(), dst_.data(), M_.data());
}
void mj_mulM_wrapper(const MjModel& m, const MjData& d, const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(res, m.nv());
CHECK_SIZE(vec, m.nv());
mj_mulM(m.get(), d.get(), res_.data(), vec_.data());
}
void mj_mulM2_wrapper(const MjModel& m, const MjData& d, const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(res, m.nv());
CHECK_SIZE(vec, m.nv());
mj_mulM2(m.get(), d.get(), res_.data(), vec_.data());
}
void mj_addM_wrapper(const MjModel& m, MjData& d, const val& dst, const val& rownnz, const val& rowadr, const val& colind)
{
UNPACK_VALUE(mjtNum, dst);
UNPACK_NULLABLE_VALUE(int, rownnz);
UNPACK_NULLABLE_VALUE(int, rowadr);
UNPACK_NULLABLE_VALUE(int, colind);
CHECK_SIZE(rownnz, m.nv());
CHECK_SIZE(rowadr, m.nv());
CHECK_SIZE(colind, m.nM());
CHECK_SIZE(dst, m.nM());
mj_addM(m.get(), d.get(), dst_.data(), rownnz_.data(), rowadr_.data(), colind_.data());
}
void mj_applyFT_wrapper(const MjModel& m, MjData& d, const NumberArray& force, const NumberArray& torque, const NumberArray& point, int body, const val& qfrc_target)
{
UNPACK_NULLABLE_ARRAY(mjtNum, force);
UNPACK_NULLABLE_ARRAY(mjtNum, torque);
UNPACK_ARRAY(mjtNum, point);
UNPACK_VALUE(mjtNum, qfrc_target);
CHECK_SIZE(qfrc_target, m.nv());
CHECK_SIZE(force, 3);
CHECK_SIZE(torque, 3);
CHECK_SIZE(point, 3);
mj_applyFT(m.get(), d.get(), force_.data(), torque_.data(), point_.data(), body, qfrc_target_.data());
}
mjtNum mj_geomDistance_wrapper(const MjModel& m, const MjData& d, int geom1, int geom2, mjtNum distmax, const val& fromto)
{
UNPACK_NULLABLE_VALUE(mjtNum, fromto);
CHECK_SIZE(fromto, 6);
return mj_geomDistance(m.get(), d.get(), geom1, geom2, distmax, fromto_.data());
}
void mj_differentiatePos_wrapper(const MjModel& m, const val& qvel, mjtNum dt, const NumberArray& qpos1, const NumberArray& qpos2)
{
UNPACK_VALUE(mjtNum, qvel);
UNPACK_ARRAY(mjtNum, qpos1);
UNPACK_ARRAY(mjtNum, qpos2);
CHECK_SIZE(qvel, m.nv());
CHECK_SIZE(qpos1, m.nq());
CHECK_SIZE(qpos2, m.nq());
mj_differentiatePos(m.get(), qvel_.data(), dt, qpos1_.data(), qpos2_.data());
}
void mj_integratePos_wrapper(const MjModel& m, const val& qpos, const NumberArray& qvel, mjtNum dt)
{
UNPACK_VALUE(mjtNum, qpos);
UNPACK_ARRAY(mjtNum, qvel);
CHECK_SIZE(qpos, m.nq());
CHECK_SIZE(qvel, m.nv());
mj_integratePos(m.get(), qpos_.data(), qvel_.data(), dt);
}
void mj_normalizeQuat_wrapper(const MjModel& m, const val& qpos)
{
UNPACK_VALUE(mjtNum, qpos);
CHECK_SIZE(qpos, m.nq());
mj_normalizeQuat(m.get(), qpos_.data());
}
void mj_multiRay_wrapper(const MjModel& m, MjData& d, const NumberArray& pnt, const NumberArray& vec, const val& geomgroup, mjtByte flg_static, int bodyexclude, const val& geomid, const val& dist, int nray, mjtNum cutoff)
{
UNPACK_ARRAY(mjtNum, pnt);
UNPACK_ARRAY(mjtNum, vec);
UNPACK_VALUE(mjtByte, geomgroup);
UNPACK_VALUE(int, geomid);
UNPACK_VALUE(mjtNum, dist);
CHECK_SIZE(dist, nray);
CHECK_SIZE(geomid, nray);
CHECK_SIZE(vec, 3 * nray);
mj_multiRay(m.get(), d.get(), pnt_.data(), vec_.data(), geomgroup_.data(), flg_static, bodyexclude, geomid_.data(), dist_.data(), nray, cutoff);
}
void mju_zero_wrapper(const val& res)
{
UNPACK_VALUE(mjtNum, res);
mju_zero(res_.data(), res_.size());
}
void mju_fill_wrapper(const val& res, mjtNum val)
{
UNPACK_VALUE(mjtNum, res);
mju_fill(res_.data(), val, res_.size());
}
void mju_copy_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_copy(res_.data(), vec_.data(), res_.size());
}
mjtNum mju_sum_wrapper(const NumberArray& vec)
{
UNPACK_ARRAY(mjtNum, vec);
return mju_sum(vec_.data(), vec_.size());
}
mjtNum mju_L1_wrapper(const NumberArray& vec)
{
UNPACK_ARRAY(mjtNum, vec);
return mju_L1(vec_.data(), vec_.size());
}
void mju_scl_wrapper(const val& res, const NumberArray& vec, mjtNum scl)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_scl(res_.data(), vec_.data(), scl, res_.size());
}
void mju_add_wrapper(const val& res, const NumberArray& vec1, const NumberArray& vec2)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec1);
UNPACK_ARRAY(mjtNum, vec2);
CHECK_SIZES(res, vec1);
CHECK_SIZES(res, vec2);
mju_add(res_.data(), vec1_.data(), vec2_.data(), res_.size());
}
void mju_sub_wrapper(const val& res, const NumberArray& vec1, const NumberArray& vec2)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec1);
UNPACK_ARRAY(mjtNum, vec2);
CHECK_SIZES(res, vec1);
CHECK_SIZES(res, vec2);
mju_sub(res_.data(), vec1_.data(), vec2_.data(), res_.size());
}
void mju_addTo_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_addTo(res_.data(), vec_.data(), res_.size());
}
void mju_subFrom_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_subFrom(res_.data(), vec_.data(), res_.size());
}
void mju_addToScl_wrapper(const val& res, const NumberArray& vec, mjtNum scl)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_addToScl(res_.data(), vec_.data(), scl, res_.size());
}
void mju_addScl_wrapper(const val& res, const NumberArray& vec1, const NumberArray& vec2, mjtNum scl)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, vec1);
UNPACK_ARRAY(mjtNum, vec2);
CHECK_SIZES(res, vec1);
CHECK_SIZES(res, vec2);
mju_addScl(res_.data(), vec1_.data(), vec2_.data(), scl, res_.size());
}
mjtNum mju_normalize_wrapper(const val& res)
{
UNPACK_VALUE(mjtNum, res);
return mju_normalize(res_.data(), res_.size());
}
mjtNum mju_norm_wrapper(const NumberArray& res)
{
UNPACK_ARRAY(mjtNum, res);
return mju_norm(res_.data(), res_.size());
}
mjtNum mju_dot_wrapper(const NumberArray& vec1, const NumberArray& vec2)
{
UNPACK_ARRAY(mjtNum, vec1);
UNPACK_ARRAY(mjtNum, vec2);
CHECK_SIZES(vec1, vec2);
return mju_dot(vec1_.data(), vec2_.data(), vec1_.size());
}
void mju_mulMatVec_wrapper(const val& res, const NumberArray& mat,
const NumberArray& vec, int nr, int nc) {
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nr);
CHECK_SIZE(vec, nc);
mju_mulMatVec(res_.data(), mat_.data(), vec_.data(), nr, nc);
}
void mju_mulMatTVec_wrapper(const val& res, const NumberArray& mat,
const NumberArray& vec, int nr, int nc) {
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nc);
CHECK_SIZE(vec, nr);
mju_mulMatTVec(res_.data(), mat_.data(), vec_.data(), nr, nc);
}
mjtNum mju_mulVecMatVec_wrapper(const NumberArray& vec1, const NumberArray& mat, const NumberArray& vec2)
{
UNPACK_ARRAY(mjtNum, vec1);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, vec2);
int64_t vec1_times_vec2 = vec1_.size() * vec2_.size();
CHECK_SIZES(vec1, vec2);
CHECK_SIZE(mat, vec1_times_vec2);
return mju_mulVecMatVec(vec1_.data(), mat_.data(), vec2_.data(), vec1_.size());
}
void mju_transpose_wrapper(const val& res, const NumberArray& mat, int nr, int nc)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nr * nc);
mju_transpose(res_.data(), mat_.data(), nr, nc);
}
void mju_symmetrize_wrapper(const val& res, const NumberArray& mat, int n)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
CHECK_SIZE(mat, n * n);
CHECK_SIZE(res, n * n);
mju_symmetrize(res_.data(), mat_.data(), n);
}
void mju_eye_wrapper(const val& mat)
{
UNPACK_VALUE(mjtNum, mat);
CHECK_PERFECT_SQUARE(mat);
mju_eye(mat_.data(), mat_sqrt);
}
void mju_mulMatMat_wrapper(const val& res, const NumberArray& mat1, const NumberArray& mat2, int r1, int c1, int c2)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat1);
UNPACK_ARRAY(mjtNum, mat2);
CHECK_SIZE(res, r1 * c2);
CHECK_SIZE(mat1, r1 * c1);
CHECK_SIZE(mat2, c1 * c2);
mju_mulMatMat(res_.data(), mat1_.data(), mat2_.data(), r1, c1, c2);
}
void mju_mulMatMatT_wrapper(const val& res, const NumberArray& mat1, const NumberArray& mat2, int r1, int c1, int r2)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat1);
UNPACK_ARRAY(mjtNum, mat2);
CHECK_SIZE(res, r1 * r2);
CHECK_SIZE(mat1, r1 * c1);
CHECK_SIZE(mat2, r2 * c1);
mju_mulMatMatT(res_.data(), mat1_.data(), mat2_.data(), r1, c1, r2);
}
void mju_mulMatTMat_wrapper(const val& res, const NumberArray& mat1, const NumberArray& mat2, int r1, int c1, int c2)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat1);
UNPACK_ARRAY(mjtNum, mat2);
CHECK_SIZE(res, c1 * c2);
CHECK_SIZE(mat1, r1 * c1);
CHECK_SIZE(mat2, r1 * c2);
mju_mulMatTMat(res_.data(), mat1_.data(), mat2_.data(), r1, c1, c2);
}
void mju_sqrMatTD_wrapper(const val& res, const NumberArray& mat, const NumberArray& diag, int nr, int nc)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, diag);
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(res, nc * nc);
CHECK_SIZE(diag, nr);
mju_sqrMatTD(res_.data(), mat_.data(), diag_.data(), nr, nc);
}
int mju_dense2sparse_wrapper(const val& res, const NumberArray& mat, int nr, int nc, const val& rownnz, const val& rowadr, const val& colind)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_VALUE(int, rownnz);
UNPACK_VALUE(int, rowadr);
UNPACK_VALUE(int, colind);
CHECK_SIZE(mat, nr * nc);
CHECK_SIZE(rownnz, nr);
CHECK_SIZE(rowadr, nr);
CHECK_SIZE(colind, res_.size());
return mju_dense2sparse(res_.data(), mat_.data(), nr, nc, rownnz_.data(), rowadr_.data(), colind_.data(), res_.size());
}
void mju_sparse2dense_wrapper(const val& res, const NumberArray& mat, int nr, int nc, const NumberArray& rownnz, const NumberArray& rowadr, const NumberArray& colind)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(int, rownnz);
UNPACK_ARRAY(int, rowadr);
UNPACK_ARRAY(int, colind);
CHECK_SIZE(res, nr * nc);
CHECK_SIZE(rownnz, nr);
CHECK_SIZE(rowadr, nr);
mju_sparse2dense(res_.data(), mat_.data(), nr, nc, rownnz_.data(), rowadr_.data(), colind_.data());
}
int mju_cholFactor_wrapper(const val& mat, mjtNum mindiag)
{
UNPACK_VALUE(mjtNum, mat);
CHECK_PERFECT_SQUARE(mat);
return mju_cholFactor(mat_.data(), mat_sqrt, mindiag);
}
void mju_cholSolve_wrapper(const val& res, const NumberArray& mat, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, vec);
CHECK_PERFECT_SQUARE(mat);
CHECK_SIZE(res, mat_sqrt);
CHECK_SIZE(vec, mat_sqrt);
mju_cholSolve(res_.data(), mat_.data(), vec_.data(), mat_sqrt);
}
int mju_cholUpdate_wrapper(const val& mat, const val& x, int flg_plus)
{
UNPACK_VALUE(mjtNum, mat);
UNPACK_VALUE(mjtNum, x);
CHECK_PERFECT_SQUARE(mat);
CHECK_SIZE(x, mat_sqrt);
return mju_cholUpdate(mat_.data(), x_.data(), mat_sqrt, flg_plus);
}
mjtNum mju_cholFactorBand_wrapper(const val& mat, int ntotal, int nband, int ndense, mjtNum diagadd, mjtNum diagmul)
{
UNPACK_VALUE(mjtNum, mat);
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
return mju_cholFactorBand(mat_.data(), ntotal, nband, ndense, diagadd, diagmul);
}
void mju_cholSolveBand_wrapper(const val& res, const NumberArray& mat, const NumberArray& vec, int ntotal, int nband, int ndense)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
CHECK_SIZE(res, ntotal);
CHECK_SIZE(vec, ntotal);
mju_cholSolveBand(res_.data(), mat_.data(), vec_.data(), ntotal, nband, ndense);
}
void mju_band2Dense_wrapper(const val& res, const NumberArray& mat, int ntotal, int nband, int ndense, mjtByte flg_sym)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
CHECK_SIZE(res, ntotal * ntotal);
mju_band2Dense(res_.data(), mat_.data(), ntotal, nband, ndense, flg_sym);
}
void mju_dense2Band_wrapper(const val& res, const NumberArray& mat, int ntotal, int nband, int ndense)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
CHECK_SIZE(mat, ntotal * ntotal);
CHECK_SIZE(res, (ntotal - ndense) * nband + ndense * ntotal);
mju_dense2Band(res_.data(), mat_.data(), ntotal, nband, ndense);
}
void mju_bandMulMatVec_wrapper(const val& res, const NumberArray& mat, const NumberArray& vec, int ntotal, int nband, int ndense, int nvec, mjtByte flg_sym)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(mjtNum, mat);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZE(mat, (ntotal - ndense) * nband + ndense * ntotal);
CHECK_SIZE(res, ntotal * nvec);
CHECK_SIZE(vec, ntotal * nvec);
mju_bandMulMatVec(res_.data(), mat_.data(), vec_.data(), ntotal, nband, ndense, nvec, flg_sym);
}
int mju_boxQP_wrapper(const val& res, const val& R, const val& index, const NumberArray& H, const NumberArray& g, const NumberArray& lower, const NumberArray& upper)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_VALUE(mjtNum, R);
UNPACK_NULLABLE_VALUE(int, index);
UNPACK_ARRAY(mjtNum, H);
UNPACK_ARRAY(mjtNum, g);
UNPACK_NULLABLE_ARRAY(mjtNum, lower);
UNPACK_NULLABLE_ARRAY(mjtNum, upper);
CHECK_SIZES(lower, res);
CHECK_SIZES(upper, res);
CHECK_SIZES(index, res);
CHECK_SIZE(R, res_.size() * (res_.size() + 7))
CHECK_PERFECT_SQUARE(H);
CHECK_SIZES(g, res);
return mju_boxQP(res_.data(), R_.data(), index_.data(), H_.data(), g_.data(), res_.size(), lower_.data(), upper_.data());
}
void mju_encodePyramid_wrapper(const val& pyramid, const NumberArray& force, const NumberArray& mu)
{
UNPACK_VALUE(mjtNum, pyramid);
UNPACK_ARRAY(mjtNum, force);
UNPACK_ARRAY(mjtNum, mu);
CHECK_SIZE(pyramid, 2 * mu_.size());
CHECK_SIZE(force, mu_.size() + 1);
mju_encodePyramid(pyramid_.data(), force_.data(), mu_.data(), mu_.size());
}
void mju_decodePyramid_wrapper(const val& force, const NumberArray& pyramid, const NumberArray& mu)
{
UNPACK_VALUE(mjtNum, force);
UNPACK_ARRAY(mjtNum, pyramid);
UNPACK_ARRAY(mjtNum, mu);
CHECK_SIZE(pyramid, 2 * mu_.size());
CHECK_SIZE(force, mu_.size() + 1);
mju_decodePyramid(force_.data(), pyramid_.data(), mu_.data(), mu_.size());
}
int mju_isZero_wrapper(const val& vec)
{
UNPACK_VALUE(mjtNum, vec);
return mju_isZero(vec_.data(), vec_.size());
}
void mju_f2n_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(float, vec);
CHECK_SIZES(res, vec);
mju_f2n(res_.data(), vec_.data(), res_.size());
}
void mju_n2f_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(float, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_n2f(res_.data(), vec_.data(), res_.size());
}
void mju_d2n_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(mjtNum, res);
UNPACK_ARRAY(double, vec);
CHECK_SIZES(res, vec);
mju_d2n(res_.data(), vec_.data(), res_.size());
}
void mju_n2d_wrapper(const val& res, const NumberArray& vec)
{
UNPACK_VALUE(double, res);
UNPACK_ARRAY(mjtNum, vec);
CHECK_SIZES(res, vec);
mju_n2d(res_.data(), vec_.data(), res_.size());
}
void mju_insertionSort_wrapper(const val& list)
{
UNPACK_VALUE(mjtNum, list);
mju_insertionSort(list_.data(), list_.size());
}
void mju_insertionSortInt_wrapper(const val& list)
{
UNPACK_VALUE(int, list);
mju_insertionSortInt(list_.data(), list_.size());
}
void mjd_transitionFD_wrapper(const MjModel& m, MjData& d, mjtNum eps, mjtByte flg_centered, const val& A, const val& B, const val& C, const val& D)
{
UNPACK_NULLABLE_VALUE(mjtNum, A);
UNPACK_NULLABLE_VALUE(mjtNum, B);
UNPACK_NULLABLE_VALUE(mjtNum, C);
UNPACK_NULLABLE_VALUE(mjtNum, D);
CHECK_SIZE(A, (2 * m.nv() + m.na()) * (2 * m.nv() + m.na()));
CHECK_SIZE(B, (2 * m.nv() + m.na()) * m.nu());
CHECK_SIZE(C, m.nsensordata() * (2 * m.nv() + m.na()));
CHECK_SIZE(D, m.nsensordata() * m.nu());
mjd_transitionFD(m.get(), d.get(), eps, flg_centered, A_.data(), B_.data(), C_.data(), D_.data());
}
void mjd_inverseFD_wrapper(const MjModel& m, MjData& d, mjtNum eps, mjtByte flg_actuation, const val& DfDq, const val& DfDv, const val& DfDa, const val& DsDq, const val& DsDv, const val& DsDa, const val& DmDq)
{
UNPACK_NULLABLE_VALUE(mjtNum, DfDq);
UNPACK_NULLABLE_VALUE(mjtNum, DfDv);
UNPACK_NULLABLE_VALUE(mjtNum, DfDa);
UNPACK_NULLABLE_VALUE(mjtNum, DsDq);
UNPACK_NULLABLE_VALUE(mjtNum, DsDv);
UNPACK_NULLABLE_VALUE(mjtNum, DsDa);
UNPACK_NULLABLE_VALUE(mjtNum, DmDq);
CHECK_SIZE(DfDq, m.nv() * m.nv());
CHECK_SIZE(DfDv, m.nv() * m.nv());
CHECK_SIZE(DfDa, m.nv() * m.nv());
CHECK_SIZE(DsDq, m.nv() * m.nsensordata());
CHECK_SIZE(DsDv, m.nv() * m.nsensordata());
CHECK_SIZE(DsDa, m.nv() * m.nsensordata());
CHECK_SIZE(DmDq, m.nv() * m.nM());
mjd_inverseFD(m.get(), d.get(), eps, flg_actuation, DfDq_.data(), DfDv_.data(), DfDa_.data(),
DsDq_.data(), DsDv_.data(), DsDa_.data(), DmDq_.data());
}
void mjd_subQuat_wrapper(const NumberArray& qa, const NumberArray& qb, const val& Da, const val& Db)
{
UNPACK_ARRAY(mjtNum, qa);
UNPACK_ARRAY(mjtNum, qb);
UNPACK_NULLABLE_VALUE(mjtNum, Da);
UNPACK_NULLABLE_VALUE(mjtNum, Db);
CHECK_SIZE(qa, 4);
CHECK_SIZE(qb, 4);
CHECK_SIZE(Da, 9);
CHECK_SIZE(Db, 9);
mjd_subQuat(qa_.data(), qb_.data(), Da_.data(), Db_.data());
}
EMSCRIPTEN_BINDINGS(mujoco_functions) {
function("parseXMLString", &parseXMLString, take_ownership());
function("error", &error_wrapper);
+2 -2
View File
@@ -224,8 +224,8 @@ class FunctionUtilsTest(absltest.TestCase):
name="my_struct",
type=ast_nodes.PointerType(ast_nodes.ValueType("mystruct")),
)
result = functions.get_params_string((param,))
self.assertEqual(result, ["Mystruct& my_struct"])
result = functions.get_param_string(param)
self.assertEqual(result, "Mystruct& my_struct")
def test_get_params_string_maybe_with_conversion_struct_ptr(self):
param = ast_nodes.FunctionParameterDecl(