From 46dc67b7eb88408b379ec1fb8db11f763ab89aef Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Tue, 29 Jul 2025 10:19:28 -0700 Subject: [PATCH] Add function mjs_sensorDim to user API. Note: change to (unreleased) contact sensor mjSpec API PiperOrigin-RevId: 788507423 Change-Id: Ie639afb7ed01dc3f7bab43a3e812e8e0b0c67d07 --- doc/APIreference/functions.rst | 9 ++++ doc/includes/references.h | 1 + include/mujoco/mujoco.h | 3 ++ python/mujoco/introspect/functions.py | 14 +++++ python/mujoco/specs.cc | 3 ++ python/mujoco/specs_test.py | 56 ++++++++++---------- src/user/user_api.cc | 73 +++++++++++++++++++++++++++ src/user/user_api.h | 2 + src/user/user_objects.cc | 46 +++-------------- src/xml/xml_native_reader.cc | 15 +++--- 10 files changed, 149 insertions(+), 73 deletions(-) diff --git a/doc/APIreference/functions.rst b/doc/APIreference/functions.rst index eea63e8d..21885e9d 100644 --- a/doc/APIreference/functions.rst +++ b/doc/APIreference/functions.rst @@ -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 diff --git a/doc/includes/references.h b/doc/includes/references.h index b554aeb1..c711b5aa 100644 --- a/doc/includes/references.h +++ b/doc/includes/references.h @@ -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); diff --git a/include/mujoco/mujoco.h b/include/mujoco/mujoco.h index 16962b02..778cfc8a 100644 --- a/include/mujoco/mujoco.h +++ b/include/mujoco/mujoco.h @@ -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. diff --git a/python/mujoco/introspect/functions.py b/python/mujoco/introspect/functions.py index 27ec5c2e..4070d123 100644 --- a/python/mujoco/introspect/functions.py +++ b/python/mujoco/introspect/functions.py @@ -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', diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 4fec07e2..3067e385 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -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) { diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index bf8acb93..4e8e2dcc 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -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("""\ @@ -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() diff --git a/src/user/user_api.cc b/src/user/user_api.cc index 95280bbd..b8890521 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -26,6 +26,7 @@ #include #include +#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) { diff --git a/src/user/user_api.h b/src/user/user_api.h index d112fc8e..f5b5c059 100644 --- a/src/user/user_api.h +++ b/src/user/user_api.h @@ -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 ----------------------------------------------- diff --git a/src/user/user_objects.cc b/src/user/user_objects.cc index 1f620d34..d2b4e26d 100644 --- a/src/user/user_objects.cc +++ b/src/user/user_objects.cc @@ -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))) { diff --git a/src/xml/xml_native_reader.cc b/src/xml/xml_native_reader.cc index e549ae6e..5e480f24 100644 --- a/src/xml/xml_native_reader.cc +++ b/src/xml/xml_native_reader.cc @@ -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; }