# 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()