Add function mjs_sensorDim to user API.

Note: change to (unreleased) contact sensor mjSpec API
PiperOrigin-RevId: 788507423
Change-Id: Ie639afb7ed01dc3f7bab43a3e812e8e0b0c67d07
This commit is contained in:
Yuval Tassa
2025-07-29 10:19:28 -07:00
committed by Copybara-Service
parent a771fc6c09
commit 46dc67b7eb
10 changed files with 149 additions and 73 deletions
+9
View File
@@ -4453,6 +4453,15 @@ Return user payload or NULL if none found.
Delete user payload.
.. _mjs_sensorDim:
`mjs_sensorDim <#mjs_sensorDim>`__
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
.. mujoco-include:: mjs_sensorDim
Return sensor dimension.
.. _ElementInitialization:
Element initialization
+1
View File
@@ -3500,6 +3500,7 @@ void mjs_setUserValueWithCleanup(mjsElement* element, const char* key,
void (*cleanup)(const void*));
const void* mjs_getUserValue(mjsElement* element, const char* key);
void mjs_deleteUserValue(mjsElement* element, const char* key);
int mjs_sensorDim(const mjsSensor* sensor);
void mjs_defaultSpec(mjSpec* spec);
void mjs_defaultOrientation(mjsOrientation* orient);
void mjs_defaultBody(mjsBody* body);
+3
View File
@@ -1683,6 +1683,9 @@ MJAPI const void* mjs_getUserValue(mjsElement* element, const char* key);
// Delete user payload.
MJAPI void mjs_deleteUserValue(mjsElement* element, const char* key);
// Return sensor dimension.
MJAPI int mjs_sensorDim(const mjsSensor* sensor);
//---------------------------------- Element initialization ---------------------------------------
// Default spec attributes.
+14
View File
@@ -10688,6 +10688,20 @@ FUNCTIONS: Mapping[str, FunctionDecl] = dict([
),
doc='Delete user payload.',
)),
('mjs_sensorDim',
FunctionDecl(
name='mjs_sensorDim',
return_type=ValueType(name='int'),
parameters=(
FunctionParameterDecl(
name='sensor',
type=PointerType(
inner_type=ValueType(name='mjsSensor', is_const=True),
),
),
),
doc='Return sensor dimension.',
)),
('mjs_defaultSpec',
FunctionDecl(
name='mjs_defaultSpec',
+3
View File
@@ -1172,6 +1172,9 @@ PYBIND11_MODULE(_specs, m) {
mjSpec.def("delete", [](MjSpec& self, raw::MjsSensor& obj) {
mjs_delete(self.ptr, obj.element);
});
mjsSensor.def("get_data_size", [](raw::MjsSensor& self) -> int {
return mjs_sensorDim(&self);
});
// ============================= MJSFLEX =====================================
mjSpec.def("delete", [](MjSpec& self, raw::MjsFlex& obj) {
+29 -27
View File
@@ -20,7 +20,7 @@ import math
import os
import textwrap
import typing
import zipfile
import zipfile # pylint: disable=unused-import
from absl.testing import absltest
from etils import epath
@@ -542,7 +542,7 @@ class SpecsTest(absltest.TestCase):
spec.to_xml()
def test_modelname_default_class(self):
XML = textwrap.dedent("""\
xml = textwrap.dedent("""\
<mujoco model="test">
<compiler angle="radian"/>
@@ -573,7 +573,7 @@ class SpecsTest(absltest.TestCase):
spec.worldbody.add_geom(main)
spec.compile()
self.assertEqual(spec.to_xml(), XML)
self.assertEqual(spec.to_xml(), xml)
spec = mujoco.MjSpec()
spec.modelname = 'test'
@@ -588,7 +588,7 @@ class SpecsTest(absltest.TestCase):
self.assertEqual(geom2.classname.name, 'main')
spec.compile()
self.assertEqual(spec.to_xml(), XML)
self.assertEqual(spec.to_xml(), xml)
spec = mujoco.MjSpec()
spec.modelname = 'test'
@@ -604,7 +604,7 @@ class SpecsTest(absltest.TestCase):
geom2.classname = main # actually redundant, since main is always applied
spec.compile()
self.assertEqual(spec.to_xml(), XML)
self.assertEqual(spec.to_xml(), xml)
# test delete default
def1 = spec.find_default('def1')
@@ -1323,12 +1323,12 @@ class SpecsTest(absltest.TestCase):
def test_bad_contact_sensor(self):
test_cases = [
dict(
expected_error='dim must be positive in sensor',
expected_error='num (intprm[2]) must be positive in sensor',
sensor_params=dict(
type=mujoco.mjtSensor.mjSENS_CONTACT,
objtype=mujoco.mjtObj.mjOBJ_GEOM,
objname='sphere1',
dim=0,
intprm=[1, 0, 0],
),
),
dict(
@@ -1337,8 +1337,7 @@ class SpecsTest(absltest.TestCase):
type=mujoco.mjtSensor.mjSENS_CONTACT,
objtype=mujoco.mjtObj.mjOBJ_GEOM,
objname='sphere1',
dim=1,
intprm=[0, 0, 0],
intprm=[0, 0, 1],
),
),
dict(
@@ -1350,8 +1349,7 @@ class SpecsTest(absltest.TestCase):
type=mujoco.mjtSensor.mjSENS_CONTACT,
objtype=mujoco.mjtObj.mjOBJ_GEOM,
objname='sphere1',
dim=1,
intprm=[1 << 10, 0, 0],
intprm=[1 << 10, 0, 1],
),
),
dict(
@@ -1363,20 +1361,7 @@ class SpecsTest(absltest.TestCase):
type=mujoco.mjtSensor.mjSENS_CONTACT,
objtype=mujoco.mjtObj.mjOBJ_GEOM,
objname='sphere1',
dim=1,
intprm=[(1 << 10) | 1, 0, 0],
),
),
dict(
expected_error=(
'dim 2 not divisible by size 3 implied by data spec (intprm[0])'
),
sensor_params=dict(
type=mujoco.mjtSensor.mjSENS_CONTACT,
objtype=mujoco.mjtObj.mjOBJ_GEOM,
objname='sphere1',
dim=2,
intprm=[2, 0, 0], # force (size 3)
intprm=[(1 << 10) | 1, 0, 1],
),
),
dict(
@@ -1385,8 +1370,7 @@ class SpecsTest(absltest.TestCase):
type=mujoco.mjtSensor.mjSENS_CONTACT,
objtype=mujoco.mjtObj.mjOBJ_GEOM,
objname='sphere1',
dim=1,
intprm=[1, 4, 0],
intprm=[1, 4, 1],
),
),
dict(
@@ -1442,6 +1426,24 @@ class SpecsTest(absltest.TestCase):
with self.assertRaisesWithPredicateMatch(ValueError, error_predicate):
spec.compile()
def test_sensor_data_size(self):
spec = mujoco.MjSpec()
quat = spec.add_sensor(
name='framequat',
type=mujoco.mjtSensor.mjSENS_FRAMEQUAT,
objtype=mujoco.mjtObj.mjOBJ_BODY,
objname='world',
)
self.assertEqual(quat.get_data_size(), 4)
clock = spec.add_sensor(
name='clock',
type=mujoco.mjtSensor.mjSENS_CLOCK,
)
self.assertEqual(clock.get_data_size(), 1)
mj_model = spec.compile()
self.assertEqual(mj_model.sensor_dim[0], 4)
self.assertEqual(mj_model.sensor_dim[1], 1)
if __name__ == '__main__':
absltest.main()
+73
View File
@@ -26,6 +26,7 @@
#include <vector>
#include <mujoco/mujoco.h>
#include "engine/engine_support.h"
#include "user/user_model.h"
#include "user/user_objects.h"
#include "user/user_cache.h"
@@ -1043,6 +1044,78 @@ void mjs_deleteUserValue(mjsElement* element, const char* key) {
// return sensor dimension
int mjs_sensorDim(const mjsSensor* sensor) {
switch (sensor->type) {
case mjSENS_TOUCH:
case mjSENS_RANGEFINDER:
case mjSENS_JOINTPOS:
case mjSENS_JOINTVEL:
case mjSENS_TENDONPOS:
case mjSENS_TENDONVEL:
case mjSENS_ACTUATORPOS:
case mjSENS_ACTUATORVEL:
case mjSENS_ACTUATORFRC:
case mjSENS_JOINTACTFRC:
case mjSENS_TENDONACTFRC:
case mjSENS_JOINTLIMITPOS:
case mjSENS_JOINTLIMITVEL:
case mjSENS_JOINTLIMITFRC:
case mjSENS_TENDONLIMITPOS:
case mjSENS_TENDONLIMITVEL:
case mjSENS_TENDONLIMITFRC:
case mjSENS_GEOMDIST:
case mjSENS_INSIDESITE:
case mjSENS_E_POTENTIAL:
case mjSENS_E_KINETIC:
case mjSENS_CLOCK:
return 1;
case mjSENS_CAMPROJECTION:
return 2;
case mjSENS_ACCELEROMETER:
case mjSENS_VELOCIMETER:
case mjSENS_GYRO:
case mjSENS_FORCE:
case mjSENS_TORQUE:
case mjSENS_MAGNETOMETER:
case mjSENS_BALLANGVEL:
case mjSENS_FRAMEPOS:
case mjSENS_FRAMEXAXIS:
case mjSENS_FRAMEYAXIS:
case mjSENS_FRAMEZAXIS:
case mjSENS_FRAMELINVEL:
case mjSENS_FRAMEANGVEL:
case mjSENS_FRAMELINACC:
case mjSENS_FRAMEANGACC:
case mjSENS_SUBTREECOM:
case mjSENS_SUBTREELINVEL:
case mjSENS_SUBTREEANGMOM:
case mjSENS_GEOMNORMAL:
return 3;
case mjSENS_GEOMFROMTO:
return 6;
case mjSENS_BALLQUAT:
case mjSENS_FRAMEQUAT:
return 4;
case mjSENS_CONTACT:
return sensor->intprm[2] * mju_condataSize(sensor->intprm[0]);
case mjSENS_USER:
return sensor->dim;
case mjSENS_PLUGIN:
return 0; // to be filled in by plugin
}
return -1;
}
// get id
int mjs_getId(mjsElement* element) {
if (!element) {
+2
View File
@@ -420,6 +420,8 @@ MJAPI const void* mjs_getUserValue(mjsElement* element, const char* key);
// Delete user payload.
MJAPI void mjs_deleteUserValue(mjsElement* element, const char* key);
// Return sensor dimension.
MJAPI int mjs_sensorDim(const mjsSensor* sensor);
//---------------------------------- Initialization -----------------------------------------------
+8 -38
View File
@@ -6801,15 +6801,12 @@ void mjCSensor::Compile(void) {
throw mjCError(this, "sensor must be attached to site");
}
// set dim and datatype
// set datatype
if (type == mjSENS_TOUCH || type == mjSENS_RANGEFINDER) {
dim = 1;
datatype = mjDATATYPE_POSITIVE;
} else if (type == mjSENS_CAMPROJECTION) {
dim = 2;
datatype = mjDATATYPE_REAL;
} else {
dim = 3;
datatype = mjDATATYPE_REAL;
}
@@ -6845,7 +6842,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 1;
datatype = mjDATATYPE_REAL;
if (type == mjSENS_JOINTPOS) {
needstage = mjSTAGE_POS;
@@ -6863,7 +6859,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 1;
datatype = mjDATATYPE_REAL;
needstage = mjSTAGE_ACC;
break;
@@ -6876,7 +6871,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 1;
datatype = mjDATATYPE_REAL;
if (type == mjSENS_TENDONPOS) {
needstage = mjSTAGE_POS;
@@ -6894,7 +6888,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 1;
datatype = mjDATATYPE_REAL;
if (type == mjSENS_ACTUATORPOS) {
needstage = mjSTAGE_POS;
@@ -6919,11 +6912,9 @@ void mjCSensor::Compile(void) {
// set
if (type == mjSENS_BALLQUAT) {
dim = 4;
datatype = mjDATATYPE_QUATERNION;
needstage = mjSTAGE_POS;
} else {
dim = 3;
datatype = mjDATATYPE_REAL;
needstage = mjSTAGE_VEL;
}
@@ -6943,7 +6934,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 1;
datatype = mjDATATYPE_REAL;
if (type == mjSENS_JOINTLIMITPOS) {
needstage = mjSTAGE_POS;
@@ -6968,7 +6958,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 1;
datatype = mjDATATYPE_REAL;
if (type == mjSENS_TENDONLIMITPOS) {
needstage = mjSTAGE_POS;
@@ -6994,13 +6983,6 @@ void mjCSensor::Compile(void) {
throw mjCError(this, "sensor must be attached to (x)body, geom, site or camera");
}
// set dim
if (type == mjSENS_FRAMEQUAT) {
dim = 4;
} else {
dim = 3;
}
// set datatype
if (type == mjSENS_FRAMEQUAT) {
datatype = mjDATATYPE_QUATERNION;
@@ -7031,7 +7013,6 @@ void mjCSensor::Compile(void) {
}
// set
dim = 3;
datatype = mjDATATYPE_REAL;
if (type == mjSENS_SUBTREECOM) {
needstage = mjSTAGE_POS;
@@ -7048,7 +7029,6 @@ void mjCSensor::Compile(void) {
if (reftype != mjOBJ_SITE) {
throw mjCError(this, "sensor must be associated with a site");
}
dim = 1;
datatype = mjDATATYPE_REAL;
needstage = mjSTAGE_POS;
break;
@@ -7076,13 +7056,10 @@ void mjCSensor::Compile(void) {
// set
needstage = mjSTAGE_POS;
if (type == mjSENS_GEOMDIST) {
dim = 1;
datatype = mjDATATYPE_POSITIVE;
} else if (type == mjSENS_GEOMNORMAL) {
dim = 3;
datatype = mjDATATYPE_AXIS;
} else {
dim = 6;
datatype = mjDATATYPE_REAL;
}
break;
@@ -7116,11 +7093,6 @@ void mjCSensor::Compile(void) {
throw mjCError(this, "subtree2 must be a child of the world");
}
// check for non-positive dim
if (dim <= 0) {
throw mjCError(this, "dim must be positive in sensor, got %d", nullptr, dim);
}
// check for dataspec correctness
int dataspec = intprm[0];
if (dataspec <= 0) {
@@ -7136,19 +7108,17 @@ void mjCSensor::Compile(void) {
"mjNCONDATA bits", nullptr, dataspec);
}
// check for dim correctness
int size = mju_condataSize(dataspec);
if (dim % size != 0) {
throw mjCError(this, "dim %d not divisible by size %d implied by data spec (intprm[0])",
nullptr, dim, size);
}
// check for reduce correctness
int reduce = intprm[1];
if (reduce < 0 || reduce > 3) {
throw mjCError(this, "unknown reduction criterion. got %d, "
"expected one of {0, 1, 2, 3}", nullptr, reduce);
}
// check for non-positive num
if (intprm[2] <= 0) {
throw mjCError(this, "num (intprm[2]) must be positive in sensor, got %d", nullptr, dim);
}
}
needstage = mjSTAGE_ACC;
@@ -7158,7 +7128,6 @@ void mjCSensor::Compile(void) {
case mjSENS_E_POTENTIAL:
case mjSENS_E_KINETIC:
case mjSENS_CLOCK:
dim = 1;
needstage = mjSTAGE_POS;
datatype = mjDATATYPE_REAL;
break;
@@ -7179,7 +7148,6 @@ void mjCSensor::Compile(void) {
break;
case mjSENS_PLUGIN:
dim = 0; // to be filled in by the plugin later
datatype = mjDATATYPE_REAL; // no noise added to plugin sensors, this attribute is unused
if (plugin_name.empty() && plugin_instance_name.empty()) {
@@ -7204,6 +7172,8 @@ void mjCSensor::Compile(void) {
throw mjCError(this, "invalid type in sensor '%s' (id = %d)", name.c_str(), id);
}
dim = mjs_sensorDim(this);
// check cutoff for incompatible data types
if (cutoff > 0 && (datatype == mjDATATYPE_QUATERNION ||
(datatype == mjDATATYPE_AXIS && type != mjSENS_GEOMNORMAL))) {
+7 -8
View File
@@ -4216,20 +4216,19 @@ void mjXReader::Sensor(XMLElement* section) {
}
sensor->intprm[0] = dataspec;
// number of contacts, sensor dim
sensor->dim = 1;
ReadAttrInt(elem, "num", &sensor->dim);
if (sensor->dim <= 0) {
throw mjXError(elem, "'num' must be positive in sensor");
}
sensor->dim *= mju_condataSize(dataspec);
// reduction type (intprm[1])
sensor->intprm[1] = 0;
if (MapValue(elem, "reduce", &n, reduce_map, reduce_sz)) {
sensor->intprm[1] = n;
}
// number of contacts (intprm[2])
sensor->intprm[2] = 1;
ReadAttrInt(elem, "num", &sensor->intprm[2]);
if (sensor->intprm[2] <= 0) {
throw mjXError(elem, "'num' must be positive in sensor");
}
// sensor type
sensor->type = mjSENS_CONTACT;
}