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:
committed by
Copybara-Service
parent
a771fc6c09
commit
46dc67b7eb
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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 -----------------------------------------------
|
||||
|
||||
|
||||
@@ -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))) {
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user