Skip to content

Commit 0e19ff4

Browse files
committed
Add relative quaternion getter
1 parent 2f0feff commit 0e19ff4

2 files changed

Lines changed: 103 additions & 4 deletions

File tree

genesis/engine/entities/rigid_entity/rigid_entity.py

Lines changed: 43 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1346,21 +1346,60 @@ def get_pos(self, envs_idx=None):
13461346
return self._solver.get_links_pos(self.base_link_idx, envs_idx)[..., 0, :]
13471347

13481348
@gs.assert_built
1349-
def get_quat(self, envs_idx=None):
1349+
def get_quat(self, envs_idx=None, *, relative=False):
13501350
"""
13511351
Returns quaternion of the entity's base link.
13521352
13531353
Parameters
13541354
----------
13551355
envs_idx : None | array_like, optional
13561356
The indices of the environments. If None, all environments will be considered. Defaults to None.
1357+
relative : bool, optional
1358+
If True, return the quaternion relative to the initial (not current!) quaternion.
1359+
The returned quaternion ``delta`` satisfies
1360+
``abs_quat == transform_quat_by_quat(init_quat, delta)``.
1361+
Equivalently, ``delta == transform_quat_by_quat(inv_quat(init_quat), abs_quat)``.
1362+
Defaults to False.
13571363
13581364
Returns
13591365
-------
13601366
quat : torch.Tensor, shape (4,) or (n_envs, 4)
1361-
The quaternion of the entity's base link.
1362-
"""
1363-
return self._solver.get_links_quat(self.base_link_idx, envs_idx)[..., 0, :]
1367+
The quaternion of the entity's base link (absolute or relative).
1368+
"""
1369+
abs_quat = self._solver.get_links_quat(self.base_link_idx, envs_idx)[..., 0, :]
1370+
if not relative:
1371+
return abs_quat
1372+
1373+
has_free_root_qpos = self.base_link.n_joints == 1 and self.base_link.joints[0].type == gs.JOINT_TYPE.FREE
1374+
if not has_free_root_qpos:
1375+
if self._solver._options.batch_links_info:
1376+
init_quat = qd_to_torch(
1377+
self._solver.links_info.quat,
1378+
envs_idx,
1379+
self.base_link_idx,
1380+
transpose=True,
1381+
copy=True,
1382+
)
1383+
if self._solver.n_envs == 0:
1384+
init_quat = init_quat[0]
1385+
else:
1386+
init_quat = init_quat[:, 0]
1387+
else:
1388+
init_quat = torch.as_tensor(self.base_link.quat, dtype=abs_quat.dtype, device=abs_quat.device)
1389+
else:
1390+
q_start = self.base_link.q_start
1391+
init_quat = qd_to_torch(
1392+
self._solver.qpos0,
1393+
envs_idx,
1394+
slice(q_start + 3, q_start + 7),
1395+
transpose=True,
1396+
copy=True,
1397+
)
1398+
if self._solver.n_envs == 0:
1399+
init_quat = init_quat[0]
1400+
1401+
init_quat = init_quat.to(dtype=abs_quat.dtype, device=abs_quat.device)
1402+
return gu.transform_quat_by_quat(gu.inv_quat(init_quat), abs_quat)
13641403

13651404
@gs.assert_built
13661405
def get_vel(self, envs_idx=None):

tests/test_rigid_physics.py

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1604,6 +1604,62 @@ def test_set_root_pose(batch_fixed_verts, relative, show_viewer, tol):
16041604
quat_ref = quat_delta
16051605
assert_allclose(quat, quat_ref, tol=tol)
16061606

1607+
if relative:
1608+
quat_rel_ref = quat_delta
1609+
else:
1610+
quat_rel_ref = gu.transform_quat_by_quat(gu.inv_quat(quat_zero), quat_delta)
1611+
assert_allclose(entity.get_quat(relative=True), quat_rel_ref, tol=tol)
1612+
# Verify get_quat(relative=False) matches get_quat() (preserves old behavior)
1613+
assert_allclose(entity.get_quat(relative=False), quat, tol=tol)
1614+
1615+
1616+
@pytest.mark.required
1617+
def test_get_quat_relative_heterogeneous_initial_quat(show_viewer, tol):
1618+
scene = gs.Scene(
1619+
rigid_options=gs.options.RigidOptions(batch_links_info=True),
1620+
show_viewer=show_viewer,
1621+
show_FPS=False,
1622+
)
1623+
box = scene.add_entity(
1624+
morph=(
1625+
gs.morphs.Box(size=(0.04, 0.04, 0.04), pos=(0.0, 0.0, 0.1), euler=(0.0, 0.0, 0.0)),
1626+
gs.morphs.Box(size=(0.04, 0.04, 0.04), pos=(0.0, 0.0, 0.1), euler=(0.0, 45.0, 0.0)),
1627+
),
1628+
)
1629+
scene.build(n_envs=4)
1630+
1631+
quat_delta = torch.tensor(
1632+
[
1633+
[0.9238795, 0.3826834, 0.0, 0.0],
1634+
[0.8660254, 0.0, 0.5, 0.0],
1635+
[0.7071068, 0.0, 0.0, 0.7071068],
1636+
[1.0, 0.0, 0.0, 0.0],
1637+
],
1638+
dtype=gs.tc_float,
1639+
device=gs.device,
1640+
)
1641+
quat_delta = quat_delta / torch.linalg.norm(quat_delta, dim=-1, keepdim=True)
1642+
1643+
box.set_quat(quat_delta, relative=True)
1644+
1645+
assert_allclose(box.get_quat(relative=True), quat_delta, tol=tol)
1646+
assert_allclose(box.get_quat(envs_idx=[2, 3], relative=True), quat_delta[2:], tol=tol)
1647+
1648+
1649+
@pytest.mark.required
1650+
def test_get_quat_relative_non_parallel(show_viewer, tol):
1651+
scene = gs.Scene(show_viewer=show_viewer, show_FPS=False)
1652+
box = scene.add_entity(gs.morphs.Box(size=(0.04, 0.04, 0.04), pos=(0.0, 0.0, 0.1), euler=(0.0, 30.0, 0.0)))
1653+
scene.build()
1654+
1655+
quat_delta = torch.tensor([0.9238795, 0.0, 0.3826834, 0.0], dtype=gs.tc_float, device=gs.device)
1656+
quat_delta = quat_delta / torch.linalg.norm(quat_delta)
1657+
1658+
box.set_quat(quat_delta, relative=True)
1659+
quat_rel = box.get_quat(relative=True)
1660+
assert quat_rel.shape == quat_delta.shape
1661+
assert_allclose(quat_rel, quat_delta, tol=tol)
1662+
16071663

16081664
@pytest.mark.required
16091665
def test_normalized_quat(show_viewer, tol):
@@ -5062,6 +5118,10 @@ def test_merge_entities(is_fixed, merge_fixed_links, show_viewer, tol, monkeypat
50625118

50635119
attach_link = franka.get_link("attachment")
50645120
assert_allclose(attach_link.get_pos(), hand.links[0].get_pos(), tol=gs.EPS)
5121+
hand_quat_rel = hand.get_quat(relative=True)
5122+
hand_init_quat = torch.as_tensor(hand.base_link.quat, dtype=gs.tc_float, device=gs.device)
5123+
hand_quat_rel_ref = gu.transform_quat_by_quat(gu.inv_quat(hand_init_quat), hand.get_quat())
5124+
assert_allclose(hand_quat_rel, hand_quat_rel_ref, tol=tol)
50655125
offset_quat = gu.transform_quat_by_quat(hand.links[0].get_quat(), gu.inv_quat(attach_link.get_quat()))
50665126
assert_allclose(gu.quat_to_xyz(offset_quat, rpy=False, degrees=True), EULER_OFFSET, tol=tol)
50675127
for link in hand.links[slice(0, None) if merge_fixed_links else slice(1, -1)]:

0 commit comments

Comments
 (0)