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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user