Add hfield. Fixes #1655 #1491 #1695

PiperOrigin-RevId: 638447127
Change-Id: Ib1e5020a8407bc100145a6b382e985c03dd4a848
This commit is contained in:
Baruch Tabanpour
2024-05-29 16:20:27 -07:00
committed by Copybara-Service
parent dec55b59fa
commit c511d02265
13 changed files with 604 additions and 106 deletions
+11 -2
View File
@@ -14,12 +14,21 @@ General
2. Added :ref:`timeconst<actuator-position-timeconst>` attribute to the :ref:`position actuator<actuator-position>`.
When set to a positive value, the actuator is made stateful with :at:`filterexact` dynamics.
MJX
^^^
3. Add height-field collision support. Fixes :github:issue:`1491`.
4. Add a pre-compiled field ``mesh_convex`` to ``mjx.Model`` so that mesh properties can be vmapped over.
Fixes :github:issue:`1655`.
5. Fix a bug in convex mesh collisions, where erroneous edge contacts were being created even though face
separating axes were found. Fixes :github:issue:`1695`.
Bug fixes
^^^^^^^^^
3. Fixed a bug the could cause collisions to be missed when :ref:`fusestatic<compiler-fusestatic>` is enabled, as is
6. Fixed a bug the could cause collisions to be missed when :ref:`fusestatic<compiler-fusestatic>` is enabled, as is
often the case for URDF imports. Fixes :github:issue:`1069`, :github:issue:`1577`.
4. Fixed a bug that was causing the visualization of SDF iterations to write outside the size of the vector storing
7. Fixed a bug that was causing the visualization of SDF iterations to write outside the size of the vector storing
them. Fixes :github:issue:`1539`.
Version 3.1.5 (May 7, 2024)
+2 -2
View File
@@ -196,7 +196,7 @@ The following features are **fully supported** in MJX:
* - :ref:`Actuator Bias <mjtBias>`
- ``NONE``, ``AFFINE``
* - :ref:`Geom <mjtGeom>`
- ``PLANE``, ``SPHERE``, ``CAPSULE``, ``BOX``, ``MESH``
- ``PLANE``, ``HFIELD``, ``SPHERE``, ``CAPSULE``, ``BOX``, ``MESH`` are fully implemented. ``ELLIPSOID`` and ``CYLINDER`` are implemented but only collide with other primitives.
* - :ref:`Constraint <mjtConstraint>`
- ``EQUALITY``, ``LIMIT_JOINT``, ``CONTACT_FRICTIONLESS``, ``CONTACT_PYRAMIDAL``
* - :ref:`Equality <mjtEq>`
@@ -223,7 +223,7 @@ The following features are **in development** and coming soon:
* - Category
- Feature
* - :ref:`Geom <mjtGeom>`
- ``SDF``, ``HFIELD``, ``ELLIPSOID``, ``CYLINDER``
- ``SDF``. Collisions between (``SPHERE``, ``BOX``, ``MESH``, ``HFIELD``) and ``CYLINDER``. Collisions between (``BOX``, ``MESH``, ``HFIELD``) and ``ELLIPSOID``.
* - :ref:`Constraint <mjtConstraint>`
- :ref:`Frictionloss <coFriction>`, ``CONTACT_ELLIPTIC``, ``FRICTION_DOF``
* - :ref:`Integrator <mjtIntegrator>`
+275 -52
View File
@@ -14,8 +14,9 @@
# ==============================================================================
"""Convex collisions."""
from collections.abc import Callable
import functools
from typing import Tuple
from typing import Tuple, Union
import jax
from jax import numpy as jp
@@ -26,16 +27,19 @@ from mujoco.mjx._src.collision_types import Collision
from mujoco.mjx._src.collision_types import ConvexInfo
from mujoco.mjx._src.collision_types import FunctionKey
from mujoco.mjx._src.collision_types import GeomInfo
from mujoco.mjx._src.collision_types import HFieldInfo
from mujoco.mjx._src.types import Data
from mujoco.mjx._src.types import GeomType
from mujoco.mjx._src.types import Model
# pylint: enable=g-importing-member
_GeomInfo = Union[GeomInfo, ConvexInfo]
def collider(ncon: int):
"""Wraps collision functions for use by collision_driver."""
def wrapper(func):
def wrapper(collision_fn):
def collide(
m: Model, d: Data, key: FunctionKey, geom: jax.Array
) -> Collision:
@@ -45,18 +49,25 @@ def collider(ncon: int):
GeomInfo(d.geom_xpos[g2], d.geom_xmat[g2], m.geom_size[g2]),
]
in_axes = [0, 0]
fn = collision_fn
for i in [0, 1]:
if key.types[i] == GeomType.BOX:
infos[i] = mesh.box(infos[i])
in_axes[i] = jax.tree_util.tree_map(lambda x: None, infos[i]).replace(
pos=0, mat=0, face=0, vert=0
pos=0, mat=0, size=0, face=0, vert=0
)
elif key.types[i] == GeomType.MESH:
infos[i] = mesh.convex(m, key.data_ids[i], infos[i])
c, cm = infos[i], m.mesh_convex[key.data_ids[i]]
infos[i] = ConvexInfo(**vars(c), **vars(cm))
in_axes[i] = jax.tree_util.tree_map(lambda x: None, infos[i]).replace(
pos=0, mat=0
pos=0, mat=0, size=0
)
dist, pos, frame = jax.vmap(func, in_axes=in_axes)(*infos)
elif key.types[i] == GeomType.HFIELD:
hfield_info = mesh.hfield(m, key.data_ids[i])
infos[i] = hfield_info.replace(pos=infos[i].pos, mat=infos[i].mat)
in_axes[i] = hfield_info.replace(pos=0, mat=0, data=None)
fn = functools.partial(fn, subgrid_size=key.subgrid_size)
dist, pos, frame = jax.vmap(fn, in_axes=in_axes)(*infos)
if ncon > 1:
return jax.tree_util.tree_map(jp.concatenate, (dist, pos, frame))
return dist, pos, frame
@@ -242,9 +253,8 @@ def plane_convex(plane: GeomInfo, convex: ConvexInfo) -> Collision:
return dist, pos, frame
@collider(ncon=1)
def sphere_convex(sphere: GeomInfo, convex: ConvexInfo) -> Collision:
"""Calculates contact between a sphere and a convex object."""
def _sphere_convex(sphere: GeomInfo, convex: ConvexInfo) -> Collision:
"""Calculates contact between a sphere and a convex mesh."""
faces = convex.face
normals = convex.face_normal
@@ -276,7 +286,7 @@ def sphere_convex(sphere: GeomInfo, convex: ConvexInfo) -> Collision:
face_normal,
)
edge_dist = jax.vmap(
lambda plane_pt, plane_norm: (pt - plane_pt).dot(plane_norm)
lambda plane_pt, plane_norm, pt=pt: (pt - plane_pt).dot(plane_norm)
)(edge_p0, side_normals)
pt_on_face = jp.all(edge_dist <= 0) # lte to handle degenerate edges
@@ -291,7 +301,8 @@ def sphere_convex(sphere: GeomInfo, convex: ConvexInfo) -> Collision:
# Get the normal, dist, and contact position.
pt_normal, d = math.normalize_with_norm(pt - sphere_pos)
# Ensure normal points towards convex centroid.
# Ensure normal points towards convex centroid. Assume convex centroid is at
# the origin.
inside = jp.dot(pt, pt_normal) > 0
sign = jp.where(inside, -1, 1)
n = jp.where(pt_on_face | (d < 1e-6), -face_normal, sign * pt_normal)
@@ -305,11 +316,17 @@ def sphere_convex(sphere: GeomInfo, convex: ConvexInfo) -> Collision:
n = convex.mat @ n
pos = convex.mat @ pos + convex.pos
return dist, pos, n
@collider(ncon=1)
def sphere_convex(sphere: GeomInfo, convex: ConvexInfo) -> Collision:
"""Calculates contact between a sphere and a convex mesh."""
dist, pos, n = _sphere_convex(sphere, convex)
return dist, pos, math.make_frame(n)
@collider(ncon=2)
def capsule_convex(cap: GeomInfo, convex: ConvexInfo) -> Collision:
def _capsule_convex(cap: GeomInfo, convex: ConvexInfo) -> Collision:
"""Calculates contacts between a capsule and a convex object."""
# Get convex transformed normals, faces, and vertices.
faces = convex.face
@@ -431,6 +448,13 @@ def capsule_convex(cap: GeomInfo, convex: ConvexInfo) -> Collision:
dist = -jp.where(
has_edge_contact, jp.array([edge_penetration, -1]), face_penetration
)
return dist, pos, n
@collider(ncon=2)
def capsule_convex(cap: GeomInfo, convex: ConvexInfo) -> Collision:
"""Calculates contacts between a capsule and a convex object."""
dist, pos, n = _capsule_convex(cap, convex)
frame = jax.vmap(math.make_frame)(n)
return dist, pos, frame
@@ -851,6 +875,7 @@ def _sat_gaussmap(
incident_face_norm,
-best_axis,
)
dist = jp.where(is_face_separating, 1.0, dist)
# Handle edge separating axes by checking all edge pairs.
a_idx = jp.tile(jp.arange(edges_a.shape[0]), reps=edges_b.shape[0])
@@ -891,8 +916,9 @@ def _sat_gaussmap(
best_edge_idx = edge_dist.argmax()
best_edge_dist = edge_dist[best_edge_idx]
is_edge_contact = jp.where(
dist.max() < 0, best_edge_dist > dist.max() - 1e-6,
(best_edge_dist < 0) & ~jp.isinf(best_edge_dist)
dist.max() < 0.0,
best_edge_dist > dist.max() - 1e-6,
(best_edge_dist < 0) & ~jp.isinf(best_edge_dist),
)
is_edge_contact = is_edge_contact & ~is_face_separating
normal = jp.where(is_edge_contact, edge_axes[best_edge_idx], normal)
@@ -911,9 +937,45 @@ def _sat_gaussmap(
return dist, pos, normal
@collider(ncon=4)
def convex_convex(c1: ConvexInfo, c2: ConvexInfo) -> Collision:
"""Calculates contacts between two convex objects."""
def _box_box(b1: ConvexInfo, b2: ConvexInfo) -> Collision:
"""Calculates contacts between two boxes."""
faces1 = b1.face
faces2 = b2.face
to_local_pos = b2.mat.T @ (b1.pos - b2.pos)
to_local_mat = b2.mat.T @ b1.mat
faces1 = to_local_pos + faces1 @ to_local_mat.T
normals1 = b1.face_normal @ to_local_mat.T
normals2 = b2.face_normal
vertices1 = to_local_pos + b1.vert @ to_local_mat.T
vertices2 = b2.vert
unique_edges1 = jp.take(vertices1, b1.edge_dir, axis=0)
unique_edges2 = jp.take(vertices2, b2.edge_dir, axis=0)
# brute-force SAT is more performant for box-box
dist, pos, normal = _sat_bruteforce(
faces1,
faces2,
vertices1,
vertices2,
normals1,
normals2,
unique_edges1,
unique_edges2,
)
# Go back to world frame.
pos = b2.pos + pos @ b2.mat.T
n = normal @ b2.mat.T
return dist, pos, n
def _convex_convex(c1: ConvexInfo, c2: ConvexInfo) -> Collision:
"""Calculates contacts between two convex meshes."""
# pad face vertices so that we can broadcast between geom1 and geom2
# face has shape (n_face, n_vert, 3)
nvert1, nvert2 = c1.face.shape[1], c2.face.shape[1]
@@ -932,6 +994,7 @@ def convex_convex(c1: ConvexInfo, c2: ConvexInfo) -> Collision:
faces1 = c1.face
faces2 = c2.face
# convert to c2 frame
to_local_pos = c2.mat.T @ (c1.pos - c2.pos)
to_local_mat = c2.mat.T @ c1.mat
@@ -942,49 +1005,209 @@ def convex_convex(c1: ConvexInfo, c2: ConvexInfo) -> Collision:
vertices1 = to_local_pos + c1.vert @ to_local_mat.T
vertices2 = c2.vert
unique_edges1 = jp.take(vertices1, c1.edge_dir, axis=0)
unique_edges2 = jp.take(vertices2, c2.edge_dir, axis=0)
edges1 = jp.take(vertices1, c1.edge, axis=0)
edges2 = jp.take(vertices2, c2.edge, axis=0)
edge_face_normals1 = c1.edge_face_normal @ to_local_mat.T
edge_face_normals2 = c2.edge_face_normal
enable_bruteforce = (
unique_edges1.shape[0] * unique_edges2.shape[0]
< edges1[0].shape[0] * edges2[0].shape[0]
dist, pos, normal = _sat_gaussmap(
to_local_pos,
faces1,
faces2,
vertices1,
vertices2,
normals1,
normals2,
edges1,
edges2,
edge_face_normals1,
edge_face_normals2,
)
if enable_bruteforce:
dist, pos, normal = _sat_bruteforce(
faces1,
faces2,
vertices1,
vertices2,
normals1,
normals2,
unique_edges1,
unique_edges2,
)
else:
dist, pos, normal = _sat_gaussmap(
to_local_pos,
faces1,
faces2,
vertices1,
vertices2,
normals1,
normals2,
edges1,
edges2,
edge_face_normals1,
edge_face_normals2,
)
# Go back to world frame.
pos = c2.pos + pos @ c2.mat.T
normal = normal @ c2.mat.T
normal = -normal if swapped else normal
frame = jax.vmap(math.make_frame)(normal)
n = normal @ c2.mat.T
n = -n if swapped else n
return dist, pos, n
@collider(ncon=4)
def box_box(b1: ConvexInfo, b2: ConvexInfo) -> Collision:
"""Calculates contacts between two boxes."""
dist, pos, n = _box_box(b1, b2)
frame = jax.vmap(math.make_frame)(n)
return dist, pos, frame
@collider(ncon=4)
def convex_convex(c1: ConvexInfo, c2: ConvexInfo) -> Collision:
"""Calculates contacts between two convex objects."""
dist, pos, n = _convex_convex(c1, c2)
frame = jax.vmap(math.make_frame)(n)
return dist, pos, frame
def _hfield_collision(
collider_fn: Callable[[_GeomInfo, _GeomInfo], Collision],
h: HFieldInfo,
obj: _GeomInfo,
obj_rbound: jax.Array,
subgrid_size: Tuple[int, int],
) -> Collision:
"""Collides an object with prisms in a height field."""
# put obj in hfield frame
obj_pos = h.mat.T @ (obj.pos - h.pos)
obj_mat = h.mat.T @ obj.mat
xmin = obj_pos[0] - obj_rbound
ymin = obj_pos[1] - obj_rbound
cmin = jp.floor((xmin + h.size[0]) / (2 * h.size[0]) * (h.ncol - 1))
cmin = cmin.astype(int)
rmin = jp.floor((ymin + h.size[1]) / (2 * h.size[1]) * (h.nrow - 1))
rmin = rmin.astype(int)
# compute real-valued grid step
dx = 2.0 * h.size[0] / (h.ncol - 1)
dy = 2.0 * h.size[1] / (h.nrow - 1)
# set zbottom value using base size
bvert = jp.array([0.0, 0.0, -h.size[3]])
bmask = jp.array([True, True, False])
# process all prisms in sub-grid
prisms = []
for r in range(subgrid_size[1]):
for c in range(subgrid_size[0]):
ri, ci = rmin + r, cmin + c
# ensure ri, ci are in the bounds of the hfield
ri = jp.clip(ri, 0, h.nrow - 2)
ci = jp.clip(ci, 0, h.ncol - 2)
p1 = [
dx * ci - h.size[0],
dy * ri - h.size[1],
h.data[ci, ri] * h.size[2],
]
p2 = [
dx * (ci + 1) - h.size[0],
dy * (ri + 1) - h.size[1],
h.data[ci + 1, ri + 1] * h.size[2],
]
p3 = [
dx * ci - h.size[0],
dy * (ri + 1) - h.size[1],
h.data[ci, ri + 1] * h.size[2],
]
top = jp.array([p1, p2, p3])
bottom = jp.array([p1, p3, p2]) * bmask + bvert
vert = jp.concatenate([bottom, top])
prisms.append(mesh.hfield_prism(vert))
p3 = p2
p2 = [
dx * (ci + 1) - h.size[0],
dy * ri - h.size[1],
h.data[ci + 1, ri] * h.size[2],
]
top = jp.array([p1, p2, p3])
bottom = jp.array([p1, p3, p2]) * bmask + bvert
vert = jp.concatenate([bottom, top])
# NB: If the order of verts is updated above, the corresponding
# hfield_prism function must be updated to ensure that all faces have the
# correct winding order.
prisms.append(mesh.hfield_prism(vert))
n_prisms = len(prisms)
prisms = jax.tree_util.tree_map(lambda *x: jp.stack(x), *prisms)
dist, pos, n = jax.vmap(collider_fn, in_axes=[None, 0])(
obj.replace(pos=obj_pos, mat=obj_mat), prisms
)
dist = dist.flatten()
pos = pos.reshape((-1, 3))
n = n.reshape((-1, 3))
n *= -1 # flip the normal since we flipped args in the call to collider_fn
# Check that we're in the half-space of the hfield norm. If not, pick the top
# face norm. This resolves issues with cracks of doom.
n_repeats = dist.shape[0] // n_prisms
top_norm = jp.repeat(prisms.face_normal[:, 1], n_repeats, axis=0)
cond = jax.vmap(jp.dot, in_axes=[0, None])(n, h.mat[2]) < 1e-6
n = jp.where(cond[:, None], top_norm, n)
return dist, pos, n
@collider(ncon=4)
def hfield_sphere(
h: HFieldInfo, s: GeomInfo, subgrid_size: Tuple[int, int]
) -> Collision:
"""Calculates contacts between a hfield and a sphere."""
rbound = jp.max(s.size)
dist, pos, n = _hfield_collision(_sphere_convex, h, s, rbound, subgrid_size)
n_mean = jp.mean(n, axis=0)
mask = dist < jp.minimum(0, dist.min() + 1e-3)
idx = _manifold_points(pos, mask, n_mean)
dist, pos, n = dist[idx], pos[idx], n[idx]
# zero out non-unique contacts
unique = jp.tril(idx == idx[:, None]).sum(axis=1) == 1
dist = jp.where(unique, dist, 1)
# back to world frame, _hfield_collision returns collision in hfield frame
pos = jax.vmap(lambda p: h.mat @ p + h.pos)(pos)
n = jax.vmap(lambda n: h.mat @ n)(n)
return dist, pos, jax.vmap(math.make_frame)(n)
@collider(ncon=4)
def hfield_capsule(
h: HFieldInfo, c: GeomInfo, subgrid_size: Tuple[int, int]
) -> Collision:
"""Calculates contacts between a hfield and a capsule."""
rbound = c.size[0] + c.size[1]
dist, pos, n = _hfield_collision(_capsule_convex, h, c, rbound, subgrid_size)
n_mean = jp.mean(n, axis=0)
mask = dist < jp.minimum(0, dist.min() + 1e-3)
idx = _manifold_points(pos, mask, n_mean)
dist, pos, n = dist[idx], pos[idx], n[idx]
# zero out non-unique contacts
unique = jp.tril(idx == idx[:, None]).sum(axis=1) == 1
dist = jp.where(unique, dist, 1)
# back to world frame, _hfield_collision returns collision in hfield frame
pos = jax.vmap(lambda p: h.mat @ p + h.pos)(pos)
n = jax.vmap(lambda n: h.mat @ n)(n)
return dist, pos, jax.vmap(math.make_frame)(n)
@collider(ncon=4)
def hfield_convex(
h: HFieldInfo, c: ConvexInfo, subgrid_size: Tuple[int, int]
) -> Collision:
"""Calculates contacts between a hfield and a capsule."""
rbound = jp.max(c.size)
dist, pos, n = _hfield_collision(_convex_convex, h, c, rbound, subgrid_size)
n_mean = jp.mean(n, axis=0)
mask = dist < jp.minimum(0, dist.min() + 1e-3)
idx = _manifold_points(pos, mask, n_mean)
dist, pos, n = dist[idx], pos[idx], n[idx]
# zero out non-unique contacts
unique = jp.tril(idx == idx[:, None]).sum(axis=1) == 1
dist = jp.where(unique, dist, 1)
# back to world frame, _hfield_collision returns collision in hfield frame
pos = jax.vmap(lambda p: h.mat @ p + h.pos)(pos)
n = jax.vmap(lambda n: h.mat @ n)(n)
return dist, pos, jax.vmap(math.make_frame)(n)
+24 -1
View File
@@ -45,8 +45,12 @@ from jax import numpy as jp
import mujoco
from mujoco.mjx._src import support
# pylint: disable=g-importing-member
from mujoco.mjx._src.collision_convex import box_box
from mujoco.mjx._src.collision_convex import capsule_convex
from mujoco.mjx._src.collision_convex import convex_convex
from mujoco.mjx._src.collision_convex import hfield_capsule
from mujoco.mjx._src.collision_convex import hfield_convex
from mujoco.mjx._src.collision_convex import hfield_sphere
from mujoco.mjx._src.collision_convex import plane_convex
from mujoco.mjx._src.collision_convex import sphere_convex
from mujoco.mjx._src.collision_primitive import capsule_capsule
@@ -78,6 +82,10 @@ _COLLISION_FUNC = {
(GeomType.PLANE, GeomType.ELLIPSOID): plane_ellipsoid,
(GeomType.PLANE, GeomType.CYLINDER): plane_cylinder,
(GeomType.PLANE, GeomType.MESH): plane_convex,
(GeomType.HFIELD, GeomType.SPHERE): hfield_sphere,
(GeomType.HFIELD, GeomType.CAPSULE): hfield_capsule,
(GeomType.HFIELD, GeomType.BOX): hfield_convex,
(GeomType.HFIELD, GeomType.MESH): hfield_convex,
(GeomType.SPHERE, GeomType.SPHERE): sphere_sphere,
(GeomType.SPHERE, GeomType.CAPSULE): sphere_capsule,
(GeomType.SPHERE, GeomType.BOX): sphere_convex,
@@ -90,7 +98,7 @@ _COLLISION_FUNC = {
(GeomType.ELLIPSOID, GeomType.ELLIPSOID): ellipsoid_ellipsoid,
(GeomType.ELLIPSOID, GeomType.CYLINDER): ellipsoid_cylinder,
(GeomType.CYLINDER, GeomType.CYLINDER): cylinder_cylinder,
(GeomType.BOX, GeomType.BOX): convex_convex,
(GeomType.BOX, GeomType.BOX): box_box,
(GeomType.BOX, GeomType.MESH): convex_convex,
(GeomType.MESH, GeomType.MESH): convex_convex,
}
@@ -210,6 +218,21 @@ def _geom_groups(
condim = max(m.geom_condim[g1], m.geom_condim[g2])
key = FunctionKey(types, data_ids, condim)
if types[0] == mujoco.mjtGeom.mjGEOM_HFIELD:
# add static grid bounds to the grouping key for hfield collisions
geom_rbound_hfield = (
m.geom_rbound_hfield if isinstance(m, Model) else m.geom_rbound
)
nrow, ncol = m.hfield_nrow[data_ids[0]], m.hfield_ncol[data_ids[0]]
xsize, ysize = m.hfield_size[data_ids[0]][:2]
xtick, ytick = (2 * xsize) / (ncol - 1), (2 * ysize) / (nrow - 1)
xbound = int(np.ceil(2 * geom_rbound_hfield[g2] / xtick)) + 1
xbound = min(xbound, ncol)
ybound = int(np.ceil(2 * geom_rbound_hfield[g2] / ytick)) + 1
ybound = min(ybound, nrow)
key = FunctionKey(types, data_ids, condim, (xbound, ybound))
groups.setdefault(key, []).append((g1, g2, ip))
return groups
@@ -721,6 +721,95 @@ class ConvexTest(absltest.TestCase):
self.assertTrue((c.dist > 0).all())
class HFieldTest(absltest.TestCase):
_HFIELD = """
<mujoco>
<asset>
<hfield name="J" size="0.9 1.1 .2 .1" nrow="8" ncol="10" elevation="
0 0 0 0 0 0 1 1 0 0
0 0 0 0 0 0 1 1 0 0
0 0 0 0 1 1 1 1 0 0
0 0 1 1 0 0 1 1 0 0
0 0 1 1 1 1 1 1 0 0
0 0 1 1 1 1 1 0 0 0
0 0 0 0 0 0 0 0 0 0
0 0 0 0 0 0 0 0 0 0"/>
</asset>
<worldbody>
<light pos="0 0 0.32"/>
<geom type="hfield" hfield="J"/>
<body pos="0 0 1">
<freejoint/>
<geom size="0.1" contype="0"/>
</body>
<body pos="0 0 0.2">
<freejoint/>
<geom type="capsule" size="0.01 0.1" contype="0"/>
</body>
<body pos="0 0 0.55">
<freejoint/>
<geom type="box" size="0.05 0.05 0.025" contype="0"/>
</body>
</worldbody>
<keyframe>
<key name="qpos1" qpos='-0.0127496 0.156995 0.118336 0.336325 -0.810241 0.442853 -0.185139 -0.19614 -0.000912274 0.112334 -0.455846 0.852431 -0.0871208 0.24078 0.124334 0.23346 0.100627 0.293376 -0.27087 0.855587 0.32944'/>
<key name="qpos2" qpos='0.0815885 -3.18397 -9.4802 0.57036 0.695092 -0.403995 -0.168298 -0.0156545 0.157173 -0.00734406 0.606091 -0.185052 0.759332 -0.147728 0.100088 -0.234066 0.224884 0.999906 0.000384 -9.62645e-05 -0.013677'/>
</keyframe>
</mujoco>
"""
def test_sphere_hfield(self):
m = mujoco.MjModel.from_xml_string(self._HFIELD)
mx = mjx.put_model(m)
d = mujoco.MjData(m)
d.qpos[:] = m.keyframe('qpos1').qpos
dx = mjx.put_data(m, d)
collision_jit_fn = jax.jit(mjx.collision)
kinematics_jit_fn = jax.jit(mjx.kinematics)
dx = kinematics_jit_fn(mx, dx)
dx = collision_jit_fn(mx, dx)
# check that all geoms are colliding with the hfield
for geom_id in [1, 2, 3]:
mask = (dx.contact.geom == np.array([0, geom_id])).all(axis=1)
c = jax.tree_util.tree_map(lambda x, m=mask: x[m], dx.contact)
self.assertTrue((c.dist < 0).any())
self.assertTrue((c.dist > -1e-3).any())
# all contact normals are roughly pointing in the right direction
self.assertTrue((c.frame[:, 0].dot(np.array([0, 0, 1])) > 0.7).all())
def test_hfield_outside(self):
"""Tests that objects outside of the hfield do not collide."""
positions = ['2.0 0', '-2.0 0', '0 -2.0', '0 2.0']
for p in positions:
xml = self._HFIELD.replace('<body pos="0 0', f'<body pos="{p}')
_, dx = _collide(xml)
self.assertTrue((dx.contact.dist >= 0).all())
def test_hfield_deep(self):
"""Tests that objects with deep penetration do not get stuck."""
m = mujoco.MjModel.from_xml_string(self._HFIELD)
mx = mjx.put_model(m)
d = mujoco.MjData(m)
d.qpos[:] = m.keyframe('qpos2').qpos
dx = mjx.put_data(m, d)
collision_jit_fn = jax.jit(mjx.collision)
kinematics_jit_fn = jax.jit(mjx.kinematics)
dx = kinematics_jit_fn(mx, dx)
dx = collision_jit_fn(mx, dx)
# check that all geoms are colliding with the hfield
for geom_id in [1, 2, 3]:
mask = (dx.contact.geom == np.array([0, geom_id])).all(axis=1)
c = jax.tree_util.tree_map(lambda x, m=mask: x[m], dx.contact)
# all contact normals are in the top half-face of the hfield
self.assertTrue((c.frame[:, 0].dot(np.array([0, 0, 1])) > 0.7).all())
class BodyPairFilterTest(absltest.TestCase):
"""Tests that certain body pairs get filtered."""
+42 -28
View File
@@ -15,31 +15,10 @@
"""Collision base types."""
import dataclasses
from typing import Tuple
from typing import Optional, Tuple
import jax
# pylint: disable=g-importing-member
from mujoco.mjx._src.dataclasses import PyTreeNode
# pylint: enable=g-importing-member
class GeomInfo(PyTreeNode):
"""Geom propertes of primitive and SDF shapes."""
pos: jax.Array
mat: jax.Array
size: jax.Array
class ConvexInfo(PyTreeNode):
"""Geom propertes of convex meshes."""
pos: jax.Array
mat: jax.Array
vert: jax.Array
face: jax.Array
face_normal: jax.Array
edge: jax.Array
edge_face_normal: jax.Array
edge_dir: jax.Array
from mujoco.mjx._src.dataclasses import PyTreeNode # pylint: disable=g-importing-member
import numpy as np
# Collision returned by collision functions:
# - distance distance between nearest points; neg: penetration
@@ -48,18 +27,53 @@ class ConvexInfo(PyTreeNode):
Collision = Tuple[jax.Array, jax.Array, jax.Array]
class GeomInfo(PyTreeNode):
"""Geom properties for primitive shapes."""
pos: jax.Array
mat: jax.Array
size: jax.Array
class ConvexInfo(PyTreeNode):
"""Geom properties for convex meshes."""
pos: jax.Array
mat: jax.Array
size: jax.Array
vert: jax.Array
face: jax.Array
face_normal: jax.Array
edge: jax.Array
edge_face_normal: jax.Array
edge_dir: Optional[jax.Array] = None
class HFieldInfo(PyTreeNode):
"""Geom properties for height fields."""
pos: jax.Array
mat: jax.Array
size: np.ndarray
nrow: int
ncol: int
data: jax.Array
@dataclasses.dataclass(frozen=True)
class FunctionKey:
"""Specifies how geom pairs group into collision_driver's function table.
Attributes:
types: geom type pair, which determines the collision function
data_ids: geom data id pair: mesh id for mesh geoms, otherwise -1.
Meshes have distinct face/vertex counts, so must occupy distinct
entries in the collision function table.
data_ids: geom data id pair: mesh id for mesh geoms, otherwise -1. Meshes
have distinct face/vertex counts, so must occupy distinct entries in the
collision function table.
condim: grouping by condim of the colliision ensures that the size of the
resulting constraint jacobian is determined at compile time.
resulting constraint jacobian is determined at compile time.
subgrid_size: the size determines the hfield subgrid to collide with
"""
types: Tuple[int, int]
data_ids: Tuple[int, int]
condim: int
subgrid_size: Tuple[int, int] = (-1, -1)
+1 -1
View File
@@ -30,7 +30,7 @@ def _jax_in_args(typ) -> bool:
return True
if dataclasses.is_dataclass(typ):
return any(_jax_in_args(f.type) for f in dataclasses.fields(typ))
if typing.get_origin(typ) in (list, dict, Union, set):
if typing.get_origin(typ) in (tuple, list, dict, Union, set):
return any(_jax_in_args(t) for t in typing.get_args(typ))
return False
+21 -3
View File
@@ -22,6 +22,7 @@ from jax import numpy as jp
import mujoco
from mujoco.mjx._src import collision_driver
from mujoco.mjx._src import constraint
from mujoco.mjx._src import mesh
from mujoco.mjx._src import support
from mujoco.mjx._src import types
import numpy as np
@@ -68,14 +69,16 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model:
if m.ntendon:
raise NotImplementedError('tendons are not supported')
mesh_geomid = set()
for g1, g2, ip in collision_driver.geom_pairs(m):
t1, t2 = m.geom_type[[g1, g2]]
# check collision function exists for type pair
if not collision_driver.has_collision_fn(t1, t2):
t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2)
raise NotImplementedError(f'({t1}, {t2}) collisions not implemented.')
# margin/gap not supported for geoms
if mujoco.mjtGeom.mjGEOM_MESH in (t1, t2):
# margin/gap not supported for meshes and height fields
no_margin = {mujoco.mjtGeom.mjGEOM_MESH, mujoco.mjtGeom.mjGEOM_HFIELD}
if no_margin.intersection({t1, t2}):
if ip != -1:
margin = m.pair_margin[ip]
else:
@@ -83,6 +86,9 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model:
if margin.any():
t1, t2 = mujoco.mjtGeom(t1), mujoco.mjtGeom(t2)
raise NotImplementedError(f'({t1}, {t2}) margin/gap not implemented.')
for t, g in [(t1, g1), (t2, g2)]:
if t == mujoco.mjtGeom.mjGEOM_MESH:
mesh_geomid.add(g)
for enum_field, enum_type, mj_type in (
(m.actuator_biastype, types.BiasType, mujoco.mjtBias),
@@ -100,12 +106,24 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model:
if not np.allclose(m.dof_frictionloss, 0):
raise NotImplementedError('dof_frictionloss is not implemented.')
fields = {f.name: getattr(m, f.name) for f in types.Model.fields()}
mjx_only = {'mesh_convex', 'geom_rbound_hfield'}
mj_field_names = {f.name for f in types.Model.fields()} - mjx_only
fields = {f: getattr(m, f) for f in mj_field_names}
fields['geom_rbound_hfield'] = fields['geom_rbound']
fields['geom_rgba'] = fields['geom_rgba'].reshape((-1, 4))
fields['mat_rgba'] = fields['mat_rgba'].reshape((-1, 4))
fields['cam_mat0'] = fields['cam_mat0'].reshape((-1, 3, 3))
fields['opt'] = _make_option(m.opt)
fields['stat'] = _make_statistic(m.stat)
# Pre-compile meshes for MJX collisions.
fields['mesh_convex'] = [None] * m.nmesh
for i in mesh_geomid:
dataid = m.geom_dataid[i]
if fields['mesh_convex'][dataid] is None:
fields['mesh_convex'][dataid] = mesh.convex(m, dataid) # pytype: disable=unsupported-operands
fields['mesh_convex'] = tuple(fields['mesh_convex'])
model = types.Model(**{k: copy.copy(v) for k, v in fields.items()})
return jax.device_put(model, device=device)
+98 -15
View File
@@ -16,14 +16,18 @@
import collections
import itertools
from typing import Tuple
from typing import Tuple, Union
import warnings
import jax
from jax import numpy as jp
import mujoco
from mujoco.mjx._src import math
# pylint: disable=g-importing-member
from mujoco.mjx._src.collision_types import ConvexInfo
from mujoco.mjx._src.collision_types import GeomInfo
from mujoco.mjx._src.collision_types import HFieldInfo
from mujoco.mjx._src.types import ConvexMesh
from mujoco.mjx._src.types import Model
# pylint: enable=g-importing-member
import numpy as np
@@ -141,7 +145,9 @@ def _convex_hull_2d(points: np.ndarray, normal: np.ndarray) -> np.ndarray:
return hull_point_idx
def _merge_coplanar(m: Model, tm: trimesh.Trimesh, meshid: int) -> np.ndarray:
def _merge_coplanar(
m: Union[mujoco.MjModel, Model], tm: trimesh.Trimesh, meshid: int
) -> np.ndarray:
"""Merges coplanar facets."""
if not tm.facets:
return tm.faces.copy() # no facets
@@ -221,6 +227,7 @@ def box(info: GeomInfo) -> ConvexInfo:
c = ConvexInfo(
info.pos,
info.mat,
info.size,
vert,
face,
face_normal,
@@ -236,22 +243,21 @@ def box(info: GeomInfo) -> ConvexInfo:
return c
def convex(m: Model, mesh_id: int, info: GeomInfo) -> ConvexInfo:
def convex(m: Union[mujoco.MjModel, Model], data_id: int) -> ConvexMesh:
"""Processes a mesh for use in convex collision algorithms.
Args:
m: an MJX model
mesh_id: the mesh id to process
info: pos, mat, size of this geom
data_id: the mesh id to process
Returns:
a convex mesh info
a convex mesh
"""
vert_beg = m.mesh_vertadr[mesh_id]
vert_end = m.mesh_vertadr[mesh_id + 1] if mesh_id < m.nmesh - 1 else None
vert_beg = m.mesh_vertadr[data_id]
vert_end = m.mesh_vertadr[data_id + 1] if data_id < m.nmesh - 1 else None
vert = m.mesh_vert[vert_beg:vert_end]
graphadr = m.mesh_graphadr[mesh_id]
graphadr = m.mesh_graphadr[data_id]
graph = m.mesh_graph[graphadr:]
graph_idx = 0
@@ -273,21 +279,98 @@ def convex(m: Model, mesh_id: int, info: GeomInfo) -> ConvexInfo:
tm_convex = trimesh.Trimesh(vertices=vert, faces=face)
vert = np.array(tm_convex.vertices)
face = _merge_coplanar(m, tm_convex, mesh_id)
face = _merge_coplanar(m, tm_convex, data_id)
face_normal = _get_face_norm(vert, face)
edge, edge_face_normal = _get_edge_normals(face, face_normal)
edge_dir = _get_unique_edge_dir(vert, face)
face = vert[face] # materialize full nface x nvert matrix
c = ConvexInfo(
info.pos,
info.mat,
c = ConvexMesh(
vert,
face,
face_normal,
edge,
edge_face_normal,
edge_dir,
)
return jax.tree_util.tree_map(jp.array, c)
def hfield_prism(vert: jax.Array) -> ConvexInfo:
"""Builds a hfield prism."""
# The first 3 vertices define the bottom triangle, and the next 3 vertices
# define the top triangle. The remaining triangles define the side of the
# prism.
face = np.array([
[0, 1, 2, 0], # bottom
[3, 4, 5, 3], # top
[0, 3, 5, 1],
[0, 2, 4, 3],
[2, 1, 5, 4],
])
edges = np.array([
# bottom
[0, 1],
[1, 2],
[0, 2],
# top
[3, 4],
[3, 5],
[4, 5],
# sides
[0, 3],
[1, 5],
[2, 4],
])
edge_face_norm = np.array([
# bottom
[0, 2],
[0, 4],
[0, 3],
# top
[1, 3],
[1, 2],
[1, 4],
# sides
[2, 3],
[2, 4],
[3, 4],
])
def get_face_norm(face):
# use ccw winding order convention, and avoid using the last vertex
edge0 = face[2, :] - face[1, :]
edge1 = face[0, :] - face[1, :]
return math.normalize(jp.cross(edge0, edge1))
centroid = jp.mean(vert, axis=0)
vert = vert - centroid
face = vert[face]
face_norm = jax.vmap(get_face_norm)(face)
c = ConvexInfo(
centroid,
jp.eye(3, dtype=float),
jp.ones(3),
vert,
face,
face_norm,
edges,
face_norm[edge_face_norm],
None,
)
return jax.tree_util.tree_map(jp.array, c)
def hfield(m: Union[mujoco.MjModel, Model], data_id: int) -> HFieldInfo:
adr = m.hfield_adr[data_id]
nrow, ncol = m.hfield_nrow[data_id], m.hfield_ncol[data_id]
h = HFieldInfo(
jp.zeros(3, dtype=float),
jp.eye(3, dtype=float),
m.hfield_size[data_id],
nrow,
ncol,
m.hfield_data[adr : adr + nrow * ncol].reshape((ncol, nrow), order='F'),
)
return h
+2
View File
@@ -197,6 +197,7 @@ def _getnum(m: Union[Model, mujoco.MjModel], obj: mujoco._enums.mjtObj) -> int:
mujoco.mjtObj.mjOBJ_SITE: m.nsite,
mujoco.mjtObj.mjOBJ_CAMERA: m.ncam,
mujoco.mjtObj.mjOBJ_MESH: m.nmesh,
mujoco.mjtObj.mjOBJ_HFIELD: m.nhfield,
mujoco.mjtObj.mjOBJ_PAIR: m.npair,
mujoco.mjtObj.mjOBJ_EQUALITY: m.neq,
mujoco.mjtObj.mjOBJ_ACTUATOR: m.nu,
@@ -218,6 +219,7 @@ def _getadr(
mujoco.mjtObj.mjOBJ_SITE: m.name_siteadr,
mujoco.mjtObj.mjOBJ_CAMERA: m.name_camadr,
mujoco.mjtObj.mjOBJ_MESH: m.name_meshadr,
mujoco.mjtObj.mjOBJ_HFIELD: m.name_hfieldadr,
mujoco.mjtObj.mjOBJ_PAIR: m.name_pairadr,
mujoco.mjtObj.mjOBJ_EQUALITY: m.name_eqadr,
mujoco.mjtObj.mjOBJ_ACTUATOR: m.name_actuatoradr,
+36 -1
View File
@@ -15,7 +15,7 @@
"""Base types used in MJX."""
import enum
from typing import Tuple
import jax
import mujoco
from mujoco.mjx._src.dataclasses import PyTreeNode # pylint: disable=g-importing-member
@@ -112,6 +112,24 @@ class GeomType(enum.IntEnum):
# unsupported: NGEOMTYPES, ARROW*, LINE, SKIN, LABEL, NONE
class ConvexMesh(PyTreeNode):
"""Geom properties for convex meshes.
Attributes:
vert: vertices of the convex mesh
face: faces of the convex mesh
face_normal: normal vectors for the faces
edge: edge indexes for all edges in the convex mesh
edge_face_normal: indexes for face normals adjacent to edges in `edge`
"""
vert: jax.Array
face: jax.Array
face_normal: jax.Array
edge: jax.Array
edge_face_normal: jax.Array
class ConeType(enum.IntEnum):
"""Type of friction cone.
@@ -322,6 +340,7 @@ class Model(PyTreeNode):
nmesh: number of meshes
nmeshvert: number of vertices in all meshes
nmeshface: number of triangular faces in all meshes
nhfield: number of heightfields
nmat: number of materials
npair: number of predefined geom pairs
nexclude: number of excluded geom pairs
@@ -396,6 +415,7 @@ class Model(PyTreeNode):
geom_solimp: constraint solver impedance: contact (ngeom, mjNIMP)
geom_size: geom-specific size parameters (ngeom, 3)
geom_rbound: radius of bounding sphere (ngeom,)
geom_rbound_hfield: static rbound for hfield grid bounds (ngeom,)
geom_pos: local position offset rel. to body (ngeom, 3)
geom_quat: local orientation offset rel. to body (ngeom, 4)
geom_friction: friction for (slide, spin, roll) (ngeom, 3)
@@ -419,6 +439,12 @@ class Model(PyTreeNode):
mesh_vert: vertex positions for all meshes (nmeshvert, 3)
mesh_face: vertex face data (nmeshface, 3)
mesh_graph: convex graph data (nmeshgraph,)
mesh_convex: pre-compiled convex mesh info for MJX (nmesh,)
hfield_size: (x, y, z_top, z_bottom) (nhfield,)
hfield_nrow: number of rows in grid (nhfield,)
hfield_ncol: number of columns in grid (nhfield,)
hfield_adr: address in hfield_data (nhfield,)
hfield_data: elevation data (nhfielddata,)
mat_rgba: rgba (nmat, 4)
pair_dim: contact dimensionality (npair,)
pair_geom1: id of geom1 (npair,)
@@ -488,6 +514,7 @@ class Model(PyTreeNode):
nmesh: int
nmeshvert: int
nmeshface: int
nhfield: int
nmat: int
npair: int
nexclude: int
@@ -561,6 +588,7 @@ class Model(PyTreeNode):
geom_solimp: jax.Array
geom_size: jax.Array
geom_rbound: jax.Array
geom_rbound_hfield: np.ndarray
geom_pos: jax.Array
geom_quat: jax.Array
geom_friction: jax.Array
@@ -584,6 +612,12 @@ class Model(PyTreeNode):
mesh_vert: np.ndarray
mesh_face: np.ndarray
mesh_graph: np.ndarray
mesh_convex: Tuple[ConvexMesh, ...]
hfield_size: np.ndarray
hfield_nrow: np.ndarray
hfield_ncol: np.ndarray
hfield_adr: np.ndarray
hfield_data: jax.Array
mat_rgba: np.ndarray
pair_dim: np.ndarray
pair_geom1: np.ndarray
@@ -632,6 +666,7 @@ class Model(PyTreeNode):
name_siteadr: np.ndarray
name_camadr: np.ndarray
name_meshadr: np.ndarray
name_hfieldadr: np.ndarray
name_pairadr: np.ndarray
name_eqadr: np.ndarray
name_actuatoradr: np.ndarray
@@ -77,6 +77,7 @@
</custom>
<asset>
<hfield name="rand_uniform" file="hfield_240_280.png" size="20.1 20.0 .05 .1"/>
<mesh file="head.stl"/>
<mesh file="powercable.stl"/>
<mesh file="handle.stl"/>
@@ -101,7 +102,8 @@
<site name="origin"/>
<light pos="0 0 1.5" dir="0 0 -1" directional="true"/>
<camera name="default" pos="0.846 -1.465 0.916" xyaxes="0.866 0.500 0.000 -0.171 0.296 0.940"/>
<geom name="floor" size="0 0 0.05" type="plane" conaffinity="1" material="groundplane"/>
<geom name="hfield_floor" type="hfield" hfield="rand_uniform" conaffinity="1"/>
<!-- <geom name="floor" pos="0 0 -0.05" size="0 0 0.05" type="plane" conaffinity="1" material="groundplane"/> -->
<body name="chassis">
<camera name="track" pos="0.846 -1.465 0.916" xyaxes="0.866 0.500 0.000 -0.171 0.296 0.940" mode="trackcom"/>
<freejoint/>
Binary file not shown.

After

Width:  |  Height:  |  Size: 170 KiB