Fix jax deprecation warning for jax.tree_map.

PiperOrigin-RevId: 626092308
Change-Id: I78145efe62ae118726c1755856d7cb9b1abd988a
This commit is contained in:
Google DeepMind
2024-04-18 11:12:41 -07:00
committed by Copybara-Service
parent 4b7b40d329
commit 8b7f1094f1
15 changed files with 52 additions and 52 deletions
+2 -2
View File
@@ -262,7 +262,7 @@ def sphere_convex(sphere: GeomInfo, convex: GeomInfo) -> Contact:
n = convex.mat @ n
pos = convex.mat @ pos + convex.pos
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
@@ -347,7 +347,7 @@ def capsule_convex(cap: GeomInfo, convex: GeomInfo) -> Contact:
degenerate_edge_dir,
edge_closest_pt,
cap_closest_pt,
) = jax.tree_map(lambda x, i=e_idx: jp.take(x, i, axis=0), res)
) = jax.tree_util.tree_map(lambda x, i=e_idx: jp.take(x, i, axis=0), res)
edge_face_normals = edge_face_normal[e_idx]
edge_voronoi_front = ((edge_face_normals @ edge_axis) < 0).all()
+7 -7
View File
@@ -216,7 +216,7 @@ def get_params(
else:
params.append(_dynamic_params(m, candidates))
params = jax.tree_map(lambda *x: jp.concatenate(x), *params)
params = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *params)
return geom1, geom2, params
@@ -232,7 +232,7 @@ def _pair_info(
d.geom_xmat[g],
m.geom_size[g],
)
in_axes = jax.tree_map(lambda x: 0, info)
in_axes = jax.tree_util.tree_map(lambda x: 0, info)
is_mesh = m.geom_convex_face[geom[0]] is not None
if is_mesh:
info = info.replace(
@@ -322,16 +322,16 @@ def _collide_geoms(
size2 = jp.max(m.geom_size[g2.geom_id], axis=-1)
dists = jax.vmap(jp.linalg.norm)(g2.pos - g1.pos) - (size1 + size2)
_, idx = jax.lax.top_k(-dists, k=n_pairs)
g1, g2, params = jax.tree_map(
g1, g2, params = jax.tree_util.tree_map(
lambda x, idx=idx: x[idx, ...], (g1, g2, params)
)
# call contact function
res = jax.vmap(fn, in_axes=in_axes)(g1, g2)
dist, pos, frame = jax.tree_map(jp.concatenate, res)
dist, pos, frame = jax.tree_util.tree_map(jp.concatenate, res)
# repeat params by the number of contacts per geom pair
geom1, geom2, params = jax.tree_map(
geom1, geom2, params = jax.tree_util.tree_map(
lambda x: jp.repeat(x, fn.ncon, axis=0), # pytype: disable=attribute-error
(g1.geom_id, g2.geom_id, params),
)
@@ -429,12 +429,12 @@ def collision(m: Model, d: Data) -> Data:
if not contacts:
raise RuntimeError('No contacts found.')
contact = jax.tree_map(lambda *x: jp.concatenate(x), *contacts)
contact = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *contacts)
max_contact_points = int(support.get_custom_numeric(m, 'max_contact_points'))
if max_contact_points > -1 and contact.dist.shape[0] > max_contact_points:
# get top-k contacts
_, idx = jax.lax.top_k(-contact.dist, k=max_contact_points)
contact = jax.tree_map(lambda x, idx=idx: jp.take(x, idx, axis=0), contact)
contact = jax.tree_util.tree_map(lambda x, idx=idx: jp.take(x, idx, axis=0), contact)
return d.replace(contact=contact)
+6 -6
View File
@@ -402,7 +402,7 @@ class CapsuleCollisionTest(parameterized.TestCase):
self.assertEqual(c.pos.shape[0], 2)
self.assertGreater(c.dist[1], 0)
# extract the contact point with penetration
c = jax.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
c = jax.tree_util.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'capsule_convex_edge', 1e-4)
@@ -437,7 +437,7 @@ class CapsuleCollisionTest(parameterized.TestCase):
self.assertEqual(c.pos.shape[0], 2)
self.assertGreater(c.dist[1], 0)
# extract the contact point with penetration
c = jax.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
c = jax.tree_util.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'edge_shallow_tip1', 1e-4)
np.testing.assert_array_almost_equal(
@@ -457,7 +457,7 @@ class CapsuleCollisionTest(parameterized.TestCase):
self.assertEqual(c.pos.shape[0], 2)
self.assertGreater(c.dist[1], 0)
# extract the contact point with penetration
c = jax.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
c = jax.tree_util.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'edge_shallow_tip2', 1e-4)
np.testing.assert_array_almost_equal(
@@ -494,7 +494,7 @@ class CylinderTest(absltest.TestCase):
d.contact.pos[:] = d.contact.pos[idx]
# extract the contact points with penetration
c = jax.tree_map(lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact)
c = jax.tree_util.tree_map(lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'cylinder_plane', 1e-5)
@@ -531,7 +531,7 @@ class ConvexTest(absltest.TestCase):
np.testing.assert_array_less(dx.contact.dist[:2], 0)
np.testing.assert_array_less(-dx.contact.dist[2:], 0)
# extract the contact points with penetration
c = jax.tree_map(lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact)
c = jax.tree_util.tree_map(lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'box_plane', 1e-5)
@@ -615,7 +615,7 @@ class ConvexTest(absltest.TestCase):
np.testing.assert_array_less(dx.contact.dist[:1], 0)
np.testing.assert_array_less(-dx.contact.dist[1:], 0)
# extract the contact point with penetration
c = jax.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
c = jax.tree_util.tree_map(lambda x: jp.take(x, 0, axis=0)[None], dx.contact)
for field in dataclasses.fields(Contact):
_assert_attr_eq(c, d.contact, field.name, 'box_box_edge', 1e-2)
+6 -6
View File
@@ -42,7 +42,7 @@ def plane_sphere(plane: GeomInfo, sphere: GeomInfo) -> Contact:
"""Calculates contact between a plane and a sphere."""
n = plane.mat[:, 2]
dist, pos = _plane_sphere(n, plane.pos, sphere.pos, sphere.size[0])
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
@@ -62,7 +62,7 @@ def plane_capsule(plane: GeomInfo, cap: GeomInfo) -> Contact:
dist = jp.expand_dims(dist, axis=0)
pos = jp.expand_dims(pos, axis=0)
contacts.append((dist, pos, frame))
return jax.tree_map(lambda *x: jp.concatenate(x), *contacts)
return jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *contacts)
def plane_ellipsoid(plane: GeomInfo, ellipsoid: GeomInfo) -> Contact:
@@ -73,7 +73,7 @@ def plane_ellipsoid(plane: GeomInfo, ellipsoid: GeomInfo) -> Contact:
pos = ellipsoid.pos + ellipsoid.mat @ (sphere_support * size)
dist = jp.dot(n, pos - plane.pos)
pos = pos - n * dist * 0.5
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
@@ -151,7 +151,7 @@ def _sphere_sphere(
def sphere_sphere(s1: GeomInfo, s2: GeomInfo) -> Contact:
"""Calculates contact between two spheres."""
dist, pos, n = _sphere_sphere(s1.pos, s1.size[0], s2.pos, s2.size[0])
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
@@ -164,7 +164,7 @@ def sphere_capsule(sphere: GeomInfo, cap: GeomInfo) -> Contact:
cap.pos - segment, cap.pos + segment, sphere.pos
)
dist, pos, n = _sphere_sphere(sphere.pos, sphere.size[0], pt, cap.size[0])
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
@@ -186,7 +186,7 @@ def capsule_capsule(cap1: GeomInfo, cap2: GeomInfo) -> Contact:
)
radius1, radius2 = cap1.size[0], cap2.size[0]
dist, pos, n = _sphere_sphere(pt1, radius1, pt2, radius2)
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
+2 -2
View File
@@ -131,7 +131,7 @@ def _optim(
def capsule_ellipsoid(c: GeomInfo, e: GeomInfo) -> Contact:
""""Calculates contact between a capsule and an ellipsoid."""
pos, dist, n = _optim(_capsule, _ellipsoid, c, e)
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
@@ -139,7 +139,7 @@ def capsule_ellipsoid(c: GeomInfo, e: GeomInfo) -> Contact:
def ellipsoid_ellipsoid(e1: GeomInfo, e2: GeomInfo) -> Contact:
""""Calculates contact between two ellipsoids."""
pos, dist, n = _optim(_ellipsoid, _ellipsoid, e1, e2)
return jax.tree_map(
return jax.tree_util.tree_map(
lambda x: jp.expand_dims(x, axis=0), (dist, pos, math.make_frame(n))
)
+4 -4
View File
@@ -111,7 +111,7 @@ def _instantiate_equality_connect(m: Model, d: Data) -> Optional[_Efc]:
return j, cpos, jp.repeat(math.norm(cpos), 3)
# concatenate to drop connect grouping dimension
j, pos, pos_norm = jax.tree_map(jp.concatenate, fn(data, id1, id2))
j, pos, pos_norm = jax.tree_util.tree_map(jp.concatenate, fn(data, id1, id2))
invweight = m.body_invweight0[id1, 0] + m.body_invweight0[id2, 0]
invweight = jp.repeat(invweight, 3)
solref = jp.tile(m.eq_solref[ids], (3, 1))
@@ -164,7 +164,7 @@ def _instantiate_equality_weld(m: Model, d: Data) -> Optional[_Efc]:
return j, pos, jp.repeat(math.norm(pos), 6)
# concatenate to drop weld grouping dimension
j, pos, pos_norm = jax.tree_map(jp.concatenate, fn(data, id1, id2))
j, pos, pos_norm = jax.tree_util.tree_map(jp.concatenate, fn(data, id1, id2))
invweight = m.body_invweight0[id1] + m.body_invweight0[id2]
invweight = jp.repeat(invweight, 3)
solref = jp.tile(m.eq_solref[ids], (6, 1))
@@ -308,7 +308,7 @@ def _instantiate_contact(m: Model, d: Data) -> Optional[_Efc]:
res = fn(d.contact)
# remove contact grouping dimension:
j, invweight, pos, solref, solimp = jax.tree_map(jp.concatenate, res)
j, invweight, pos, solref, solimp = jax.tree_util.tree_map(jp.concatenate, res)
frictionloss = jp.zeros_like(pos)
return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss)
@@ -366,7 +366,7 @@ def make_constraint(m: Model, d: Data) -> Data:
d = d.replace(efc_D=z, efc_aref=z, efc_frictionloss=z)
return d
efc = jax.tree_map(lambda *x: jp.concatenate(x), *efcs)
efc = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *efcs)
@jax.vmap
def fn(efc):
+1 -1
View File
@@ -284,7 +284,7 @@ def device_get_into(result, value):
)
for i in range(batch_size):
value_i = jax.tree_map(lambda x, i=i: x[i], value)
value_i = jax.tree_util.tree_map(lambda x, i=i: x[i], value)
device_get_into(result[i], value_i)
else:
+2 -2
View File
@@ -92,7 +92,7 @@ class DeviceTest(parameterized.TestCase):
# create mjx_data and batch it
dx = mjx.make_data(mx)
dx = jax.tree_map(
dx = jax.tree_util.tree_map(
lambda x: jp.repeat(x, batch_size).reshape((batch_size,) + x.shape),
dx,
)
@@ -105,7 +105,7 @@ class DeviceTest(parameterized.TestCase):
device.device_get_into(ds, dx)
dx = jax.device_get(dx) # faster indexing for testing
for i in range(batch_size):
_assert_eq(self, jax.tree_map(lambda x, i=i: x[i], dx), ds[i])
_assert_eq(self, jax.tree_util.tree_map(lambda x, i=i: x[i], dx), ds[i])
class ValidateInputTest(absltest.TestCase):
+2 -2
View File
@@ -310,7 +310,7 @@ def rungekutta4(m: Model, d: Data) -> Data:
kqvel = d.qvel # intermediate RK solution
# RK solutions sum
qvel, qacc, act_dot = jax.tree_map(
qvel, qacc, act_dot = jax.tree_util.tree_map(
lambda k: B[0] * k, (kqvel, d.qacc, d.act_dot)
)
integrate_fn = lambda *args: _integrate_pos(*args, dt=m.opt.timestep)
@@ -318,7 +318,7 @@ def rungekutta4(m: Model, d: Data) -> Data:
def f(carry, x):
qvel, qacc, act_dot, kqvel, d = carry
a, b, t = x # tableau numbers
dqvel, dqacc, dact_dot = jax.tree_map(
dqvel, dqacc, dact_dot = jax.tree_util.tree_map(
lambda k: a * k, (kqvel, d.qacc, d.act_dot)
)
# get intermediate RK solutions
+2 -2
View File
@@ -295,7 +295,7 @@ def get_data_into(
j = m.dof_parentid[j]
for i in range(batch_size):
d_i = jax.tree_map(lambda x, i=i: x[i], d) if batched else d
d_i = jax.tree_util.tree_map(lambda x, i=i: x[i], d) if batched else d
result_i = result[i] if batched else result
ncon = (d_i.contact.dist <= 0).sum()
efc_active = (d_i.efc_J != 0).any(axis=1)
@@ -353,7 +353,7 @@ def _put_contact(
pad_fn = lambda x: np.concatenate(
(x, np.zeros((pad_size,) + x.shape[1:], dtype=x.dtype))
)
fields = jax.tree_map(pad_fn, fields)
fields = jax.tree_util.tree_map(pad_fn, fields)
fields['dist'][-pad_size:] = np.inf
fields = jax.device_put(fields, device=device)
+1 -1
View File
@@ -435,7 +435,7 @@ class DataIOTest(parameterized.TestCase):
mujoco.mj_step(m, d, 2)
dx = mjx.put_data(m, d)
# second data in batch has contact dist > 0, disables contact
dx_b = jax.tree_map(lambda x: jp.stack((x, x + 0.05)), dx)
dx_b = jax.tree_util.tree_map(lambda x: jp.stack((x, x + 0.05)), dx)
ds = mjx.get_data(m, dx_b)
self.assertLen(ds, 2)
np.testing.assert_allclose(ds[0].qpos, d.qpos)
+9 -9
View File
@@ -62,7 +62,7 @@ def _take(obj: Y, idx: np.ndarray) -> Y:
x = x.take(jp.array(idx), axis=0, mode='wrap')
return x
return jax.tree_map(take, obj)
return jax.tree_util.tree_map(take, obj)
def _q_bodyid(m: Model) -> np.ndarray:
@@ -120,7 +120,7 @@ def _nvmap(f: Callable[..., Y], *args) -> Y:
args = [a if n is None else None for n, a in zip(np_args, args)]
# remove empty args that we should not vmap over
args = jax.tree_map(lambda a: a if a.shape[0] else None, args)
args = jax.tree_util.tree_map(lambda a: a if a.shape[0] else None, args)
in_axes = [None if a is None else 0 for a in args]
def outer_f(*args, np_args=np_args):
@@ -322,7 +322,7 @@ def flat(
[v if typ in flat_ else jp.concatenate(v) for v, typ in zip(y, out_types)]
for y in ys
]
ys = jax.tree_map(lambda *x: jp.concatenate(x), *ys)
ys = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *ys)
# put concatenated results back in order
reordered_ys = []
@@ -465,15 +465,15 @@ def body_tree(
def index_sum(x, i=id_map, s=body_ids.size):
return jax.ops.segment_sum(x, i, s)
y = jax.tree_map(index_sum, y)
carry = y if carry is None else jax.tree_map(jp.add, carry, y)
y = jax.tree_util.tree_map(index_sum, y)
carry = y if carry is None else jax.tree_util.tree_map(jp.add, carry, y)
elif key in key_parents:
ys = [key_y[p] for p in key_parents[key]]
y = jax.tree_map(lambda *x: jp.concatenate(x), *ys)
y = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *ys)
body_ids = np.concatenate([key_body_ids[p] for p in key_parents[key]])
parent_ids = m.body_parentid[key_body_ids[key]]
take_fn = lambda x, i=_index(body_ids, parent_ids): _take(x, i)
carry = jax.tree_map(take_fn, y)
carry = jax.tree_util.tree_map(take_fn, y)
f_args = [_take(arg, ids) for arg, ids in zip(args, key_in_take[key])]
key_y[key] = _nvmap(f, carry, *f_args)
@@ -488,8 +488,8 @@ def body_tree(
if len(out_types) > 1:
y_typ = [y_[i] for y_ in y_typ]
if typ != 'b':
y_typ = jax.tree_map(jp.concatenate, y_typ)
y_typ = jax.tree_map(lambda *x: jp.concatenate(x), *y_typ)
y_typ = jax.tree_util.tree_map(jp.concatenate, y_typ)
y_typ = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *y_typ)
y_take = np.argsort(np.concatenate([key_y_take[key][i] for key in keys]))
_check_output(y_typ, y_take, typ, i)
y.append(_take(y_typ, y_take))
+6 -6
View File
@@ -283,14 +283,14 @@ def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
# 1) they are not correctly at a bracket boundary (e.g. lo.deriv_0 > 0), OR
# 2) if moving to next or mid narrows the bracket
swap_lo_next = (lo.deriv_0 > 0) | (lo.deriv_0 < lo_next.deriv_0)
lo = jax.tree_map(lambda x, y: jp.where(swap_lo_next, y, x), lo, lo_next)
lo = jax.tree_util.tree_map(lambda x, y: jp.where(swap_lo_next, y, x), lo, lo_next)
swap_lo_mid = (mid.deriv_0 < 0) & (lo.deriv_0 < mid.deriv_0)
lo = jax.tree_map(lambda x, y: jp.where(swap_lo_mid, y, x), lo, mid)
lo = jax.tree_util.tree_map(lambda x, y: jp.where(swap_lo_mid, y, x), lo, mid)
swap_hi_next = (hi.deriv_0 < 0) | (hi.deriv_0 > hi_next.deriv_0)
hi = jax.tree_map(lambda x, y: jp.where(swap_hi_next, y, x), hi, hi_next)
hi = jax.tree_util.tree_map(lambda x, y: jp.where(swap_hi_next, y, x), hi, hi_next)
swap_hi_mid = (mid.deriv_0 > 0) & (hi.deriv_0 > mid.deriv_0)
hi = jax.tree_map(lambda x, y: jp.where(swap_hi_mid, y, x), hi, mid)
hi = jax.tree_util.tree_map(lambda x, y: jp.where(swap_hi_mid, y, x), hi, mid)
swap = swap_lo_next | swap_lo_mid | swap_hi_next | swap_hi_mid
@@ -302,8 +302,8 @@ def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
p0 = point_fn(jp.array(0.0))
lo = point_fn(p0.alpha - p0.deriv_0 / p0.deriv_1)
lesser_fn = lambda x, y: jp.where(lo.deriv_0 < p0.deriv_0, x, y)
hi = jax.tree_map(lesser_fn, p0, lo)
lo = jax.tree_map(lesser_fn, lo, p0)
hi = jax.tree_util.tree_map(lesser_fn, p0, lo)
lo = jax.tree_util.tree_map(lesser_fn, lo, p0)
ls_ctx = _LSContext(lo=lo, hi=hi, swap=jp.array(True), ls_iter=0)
ls_ctx = _while_loop_scan(cond, body, ls_ctx, m.opt.ls_iterations)
@@ -79,7 +79,7 @@ class CollisionDriverIntegrationTest(parameterized.TestCase):
self.assertSequenceEqual(set(idx_mjx), set(idx_mj))
idx = sorted(range(len(idx_mj)), key=lambda x: idx_mj.index(idx_mjx[x]))
mjx_contact = jax.tree_map(
mjx_contact = jax.tree_util.tree_map(
lambda x: x.take(np.array(idx), axis=0), dx.contact
)
for field in dataclasses.fields(Contact):
+1 -1
View File
@@ -829,7 +829,7 @@
"\n",
" friction, gain, bias = rand(rng)\n",
"\n",
" in_axes = jax.tree_map(lambda x: None, sys)\n",
" in_axes = jax.tree_util.tree_map(lambda x: None, sys)\n",
" in_axes = in_axes.tree_replace({\n",
" 'geom_friction': 0,\n",
" 'actuator_gainprm': 0,\n",