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:
committed by
Copybara-Service
parent
3f69deb328
commit
a36f452250
+755
-812
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user