# Copyright 2023 DeepMind Technologies Limited # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. # ============================================================================== """Tests for scan functions.""" from absl.testing import absltest from jax import numpy as jp import mujoco from mujoco import mjx # pylint: disable=g-importing-member from mujoco.mjx._src import scan from mujoco.mjx._src.types import JointType # pylint: enable=g-importing-member import numpy as np class ScanTest(absltest.TestCase): _MULTI_DOF_XML = """ """ def test_flat_empty(self): """Test scanning over just world body.""" m = mujoco.MjModel.from_xml_string(""" """) m = mjx.device_put(m) def fn(body_id): return body_id + 1 b_in = jp.array([1]) b_expect = jp.array([2]) b_out = scan.flat(m, fn, 'b', 'b', b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) def test_flat_joints(self): """Tests scanning over bodies with joints of different types.""" m = mujoco.MjModel.from_xml_string(self._MULTI_DOF_XML) m = mjx.device_put(m) # we will test two functions: # 1) j_fn receives jnt_types as a jp array # 2) s_fn receives jnt_types as a static np array and can switch on it j_fn = lambda jnt_pos, val: val + jp.sum(jnt_pos) s_fn = lambda jnt_types, val: val + sum(jnt_types) b_in = jp.array([[0, 0], [1, 1], [2, 2], [3, 3]]) b_expect = jp.array([[0, 0], [1, 1], [3, 3], [8, 8]]) b_out = scan.flat(m, j_fn, 'jb', 'b', m.jnt_pos, b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) b_out = scan.flat(m, s_fn, 'jb', 'b', m.jnt_type, b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) # None should be omitted from the results def no_free(jnt_types, val): if tuple(jnt_types) == (JointType.FREE,): return None return val + sum(jnt_types) b_expect = jp.array([[0, 0], [3, 3], [8, 8]]) b_out = scan.flat(m, no_free, 'jb', 'b', m.jnt_type, b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) # we should not call functions for which we know we will discard the results def no_world(jnt_types, val): if jnt_types.size == 0: self.fail('world has no dofs, should not be called') return val + sum(jnt_types) v_in = jp.ones((m.nv, 1)) scan.flat(m, no_world, 'jv', 'v', m.jnt_type, v_in) def test_body_tree(self): """Tests tree scanning over bodies with different joint counts.""" m = mujoco.MjModel.from_xml_string(self._MULTI_DOF_XML) m = mjx.device_put(m) # we will test two functions: # 1) j_fn receives jnt_pos which is a jp array # 2) s_fn receives jnt_types which is a static np array def j_fn(carry, jnt_pos, val): carry = jp.zeros_like(val) if carry is None else carry return carry + val + jp.sum(jnt_pos) def s_fn(carry, jnt_types, val): carry = jp.zeros_like(val) if carry is None else carry return carry + val + sum(jnt_types) b_in = jp.array([[0, 0], [1, 1], [2, 2], [3, 3]]) b_expect = jp.array([[0, 0], [1, 1], [4, 4], [9, 9]]) b_out = scan.body_tree(m, j_fn, 'jb', 'b', m.jnt_pos, b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) b_out = scan.body_tree(m, s_fn, 'jb', 'b', m.jnt_type, b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) # and reverse too: b_expect = jp.array([[12, 12], [12, 12], [3, 3], [8, 8]]) b_out = scan.body_tree(m, j_fn, 'jb', 'b', m.jnt_pos, b_in, reverse=True) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) b_out = scan.body_tree(m, s_fn, 'jb', 'b', m.jnt_type, b_in, reverse=True) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) # None should be omitted from the results def no_free(carry, jnt_types, val): if tuple(jnt_types) == (JointType.FREE,): return None carry = jp.zeros_like(val) if carry is None else carry return carry + val + sum(jnt_types) b_expect = jp.array([[0, 0], [3, 3], [8, 8]]) b_out = scan.body_tree(m, no_free, 'jb', 'b', m.jnt_type, b_in) np.testing.assert_equal(np.array(b_out), np.array(b_expect)) _MULTI_ACT_XML = """ """ def testscan_actuators(self): """Tests scanning over actuators.""" m = mujoco.MjModel.from_xml_string(self._MULTI_ACT_XML) m = mjx.device_put(m) fn = lambda *args: args args = ( m.actuator_gear, m.jnt_type, jp.arange(m.nq), jp.arange(m.nv), jp.array([1.4, 1.1]), ) gear, jnt_typ, qadr, vadr, act = scan.flat( m, fn, 'ujqva', 'ujqva', *args, group_by='u' ) actuator_trnid = m.actuator_trnid[:, 0] np.testing.assert_array_equal(gear, m.actuator_gear) np.testing.assert_array_equal(jnt_typ, m.jnt_type[actuator_trnid]) np.testing.assert_array_equal(act, jp.array([1.4, 1.1])) expected_vadr = np.concatenate( [np.nonzero(m.dof_jntid == trnid)[0] for trnid in actuator_trnid] ) np.testing.assert_array_equal(vadr, expected_vadr) expected_qadr = np.concatenate( [np.nonzero(scan._q_jointid(m) == i)[0] for i in actuator_trnid] ) np.testing.assert_array_equal(qadr, expected_qadr) if __name__ == '__main__': absltest.main()