Introduce private header engine_inline.h exploiting restrict and avoiding loops and copies in some commonly used utility functions.

PiperOrigin-RevId: 843635464
Change-Id: I3b0553eea98424ccc7def77e3769e2994f3e9014
This commit is contained in:
Yuval Tassa
2025-12-12 04:48:21 -08:00
committed by Copybara-Service
parent a0a56065e0
commit 600f0f20bc
11 changed files with 964 additions and 487 deletions
+50 -49
View File
@@ -17,6 +17,7 @@
#include <math.h>
#include <mujoco/mjmodel.h>
#include "engine/engine_inline.h"
#include "engine/engine_util_blas.h"
#include "engine/engine_util_errmem.h"
@@ -32,7 +33,7 @@ void mju_rotVecQuat(mjtNum res[3], const mjtNum vec[3], const mjtNum quat[4]) {
// null quat: copy vec
else if (quat[0] == 1 && quat[1] == 0 && quat[2] == 0 && quat[3] == 0) {
mju_copy3(res, vec);
mji_copy3(res, vec);
}
// regular processing
@@ -124,7 +125,7 @@ void mju_quat2Vel(mjtNum res[3], const mjtNum quat[4], mjtNum dt) {
}
speed /= dt;
mju_scl3(res, axis, speed);
mji_scl3(res, axis, speed);
}
@@ -132,11 +133,11 @@ void mju_quat2Vel(mjtNum res[3], const mjtNum quat[4], mjtNum dt) {
void mju_subQuat(mjtNum res[3], const mjtNum qa[4], const mjtNum qb[4]) {
// qdif = neg(qb)*qa
mjtNum qneg[4], qdif[4];
mju_negQuat(qneg, qb);
mju_mulQuat(qdif, qneg, qa);
mji_negQuat(qneg, qb);
mji_mulQuat(qdif, qneg, qa);
// convert to 3D velocity
mju_quat2Vel(res, qdif, 1);
mji_quat2Vel(res, qdif, 1);
}
@@ -157,16 +158,16 @@ void mju_quat2Mat(mjtNum res[9], const mjtNum quat[4]) {
// regular processing
else {
const mjtNum q00 = quat[0]*quat[0];
const mjtNum q01 = quat[0]*quat[1];
const mjtNum q02 = quat[0]*quat[2];
const mjtNum q03 = quat[0]*quat[3];
const mjtNum q11 = quat[1]*quat[1];
const mjtNum q12 = quat[1]*quat[2];
const mjtNum q13 = quat[1]*quat[3];
const mjtNum q22 = quat[2]*quat[2];
const mjtNum q23 = quat[2]*quat[3];
const mjtNum q33 = quat[3]*quat[3];
mjtNum q00 = quat[0]*quat[0];
mjtNum q01 = quat[0]*quat[1];
mjtNum q02 = quat[0]*quat[2];
mjtNum q03 = quat[0]*quat[3];
mjtNum q11 = quat[1]*quat[1];
mjtNum q12 = quat[1]*quat[2];
mjtNum q13 = quat[1]*quat[3];
mjtNum q22 = quat[2]*quat[2];
mjtNum q23 = quat[2]*quat[3];
mjtNum q33 = quat[3]*quat[3];
res[0] = q00 + q11 - q22 - q33;
res[4] = q00 - q11 + q22 - q33;
@@ -234,9 +235,9 @@ void mju_quatIntegrate(mjtNum quat[4], const mjtNum vel[3], mjtNum scale) {
mjtNum angle, tmp[4], qrot[4];
// form local rotation quaternion, apply
mju_copy3(tmp, vel);
mji_copy3(tmp, vel);
angle = scale * mju_normalize3(tmp);
mju_axisAngle2Quat(qrot, tmp, angle);
mji_axisAngle2Quat(qrot, tmp, angle);
mju_normalize4(quat);
mju_mulQuat(quat, quat, qrot);
}
@@ -256,7 +257,7 @@ void mju_quatZ2Vec(mjtNum quat[4], const mjtNum vec[3]) {
}
// compute angle and axis
mju_cross(axis, z, vn);
mji_cross(axis, z, vn);
a = mju_normalize3(axis);
// almost parallel
@@ -272,7 +273,7 @@ void mju_quatZ2Vec(mjtNum quat[4], const mjtNum vec[3]) {
// make quaternion from angle and axis
a = mju_atan2(a, mju_dot3(vn, z));
mju_axisAngle2Quat(quat, axis, a);
mji_axisAngle2Quat(quat, axis, a);
}
@@ -295,11 +296,11 @@ int mju_mat2Rot(mjtNum quat[4], const mjtNum mat[9]) {
mjtNum col2_rot[3] = {rot[1], rot[4], rot[7]};
mjtNum col3_rot[3] = {rot[2], rot[5], rot[8]};
mjtNum omega[3], vec1[3], vec2[3], vec3[3];
mju_cross(vec1, col1_rot, col1_mat);
mju_cross(vec2, col2_rot, col2_mat);
mju_cross(vec3, col3_rot, col3_mat);
mju_add3(omega, vec1, vec2);
mju_addTo3(omega, vec3);
mji_cross(vec1, col1_rot, col1_mat);
mji_cross(vec2, col2_rot, col2_mat);
mji_cross(vec3, col3_rot, col3_mat);
mji_add3(omega, vec1, vec2);
mji_addTo3(omega, vec3);
mju_scl3(omega, omega, 1.0 / (mju_abs(mju_dot3(col1_rot, col1_mat) +
mju_dot3(col2_rot, col2_mat) +
mju_dot3(col3_rot, col3_mat)) + mjMINVAL));
@@ -308,7 +309,7 @@ int mju_mat2Rot(mjtNum quat[4], const mjtNum mat[9]) {
break;
}
mjtNum qrot[4];
mju_axisAngle2Quat(qrot, omega, w);
mji_axisAngle2Quat(qrot, omega, w);
mju_mulQuat(quat, qrot, quat);
mju_normalize4(quat);
}
@@ -323,22 +324,22 @@ void mju_mulPose(mjtNum posres[3], mjtNum quatres[4],
const mjtNum pos1[3], const mjtNum quat1[4],
const mjtNum pos2[3], const mjtNum quat2[4]) {
// quatres = quat1*quat2
mju_mulQuat(quatres, quat1, quat2);
mji_mulQuat(quatres, quat1, quat2);
mju_normalize4(quatres);
// posres = quat1*pos2 + pos1
mju_rotVecQuat(posres, pos2, quat1);
mju_addTo3(posres, pos1);
mji_rotVecQuat(posres, pos2, quat1);
mji_addTo3(posres, pos1);
}
// negate pose
void mju_negPose(mjtNum posres[3], mjtNum quatres[4], const mjtNum pos[3], const mjtNum quat[4]) {
// qres = neg(quat)
mju_negQuat(quatres, quat);
mji_negQuat(quatres, quat);
// pres = -neg(quat)*pos
mju_rotVecQuat(posres, pos, quatres);
mji_rotVecQuat(posres, pos, quatres);
mju_scl3(posres, posres, -1);
}
@@ -346,8 +347,8 @@ void mju_negPose(mjtNum posres[3], mjtNum quatres[4], const mjtNum pos[3], const
// transform vector by pose
void mju_trnVecPose(mjtNum res[3], const mjtNum pos[3], const mjtNum quat[4], const mjtNum vec[3]) {
// res = quat*vec + pos
mju_rotVecQuat(res, vec, quat);
mju_addTo3(res, pos);
mji_rotVecQuat(res, vec, quat);
mji_addTo3(res, pos);
}
@@ -397,7 +398,7 @@ void mju_crossForce(mjtNum res[6], const mjtNum vel[6], const mjtNum f[6]) {
// express inertia in com-based frame
void mju_inertCom(mjtNum res[10], const mjtNum inert[3], const mjtNum mat[9],
void mju_inertCom(mjtNum* restrict res, const mjtNum inert[3], const mjtNum mat[9],
const mjtNum dif[3], mjtNum mass) {
// tmp = diag(inert) * mat' (mat is local-to-global rotation)
mjtNum tmp[9] = {mat[0]*inert[0], mat[3]*inert[0], mat[6]*inert[0],
@@ -442,17 +443,17 @@ void mju_mulInertVec(mjtNum* restrict res, const mjtNum i[10], const mjtNum v[6]
// express motion axis in com-based frame
void mju_dofCom(mjtNum res[6], const mjtNum axis[3], const mjtNum offset[3]) {
void mju_dofCom(mjtNum* restrict res, const mjtNum axis[3], const mjtNum offset[3]) {
// hinge
if (offset) {
mju_copy3(res, axis);
mju_cross(res+3, axis, offset);
mji_copy3(res, axis);
mji_cross(res+3, axis, offset);
}
// slide
else {
mju_zero3(res);
mju_copy3(res+3, axis);
mji_copy3(res+3, axis);
}
}
@@ -481,24 +482,24 @@ void mju_transformSpatial(mjtNum res[6], const mjtNum vec[6], int flg_force,
// apply translation
mju_copy(tran, vec, 6);
mju_sub3(dif, newpos, oldpos);
mji_sub3(dif, newpos, oldpos);
if (flg_force) {
mju_cross(cros, dif, vec+3);
mju_sub3(tran, vec, cros);
mji_cross(cros, dif, vec+3);
mji_sub3(tran, vec, cros);
} else {
mju_cross(cros, dif, vec);
mju_sub3(tran+3, vec+3, cros);
mji_cross(cros, dif, vec);
mji_sub3(tran+3, vec+3, cros);
}
// if provided, apply old -> new rotation
if (rotnew2old) {
mju_mulMatTVec3(res, rotnew2old, tran);
mju_mulMatTVec3(res+3, rotnew2old, tran+3);
mji_mulMatTVec3(res, rotnew2old, tran);
mji_mulMatTVec3(res+3, rotnew2old, tran+3);
}
// otherwise copy
else {
mju_copy(res, tran, 6);
mji_copy6(res, tran);
}
}
@@ -524,12 +525,12 @@ void mju_makeFrame(mjtNum frame[9]) {
}
// make yaxis orthogonal to xaxis
mju_scl3(tmp, frame, mju_dot3(frame, frame+3));
mju_subFrom3(frame+3, tmp);
mji_scl3(tmp, frame, mju_dot3(frame, frame+3));
mji_subFrom3(frame+3, tmp);
mju_normalize3(frame+3);
// zaxis = cross(xaxis, yaxis)
mju_cross(frame+6, frame, frame+3);
mji_cross(frame+6, frame, frame+3);
}
@@ -566,5 +567,5 @@ void mju_euler2Quat(mjtNum quat[4], const mjtNum euler[3], const char* seq) {
}
}
mju_copy4(quat, tmp);
mji_copy4(quat, tmp);
}