From 8b7f1094f199c629163ccef72d52283ae5886bdd Mon Sep 17 00:00:00 2001 From: Google DeepMind Date: Thu, 18 Apr 2024 11:12:41 -0700 Subject: [PATCH] Fix jax deprecation warning for `jax.tree_map`. PiperOrigin-RevId: 626092308 Change-Id: I78145efe62ae118726c1755856d7cb9b1abd988a --- mjx/mujoco/mjx/_src/collision_convex.py | 4 ++-- mjx/mujoco/mjx/_src/collision_driver.py | 14 +++++++------- mjx/mujoco/mjx/_src/collision_driver_test.py | 12 ++++++------ mjx/mujoco/mjx/_src/collision_primitive.py | 12 ++++++------ mjx/mujoco/mjx/_src/collision_sdf.py | 4 ++-- mjx/mujoco/mjx/_src/constraint.py | 8 ++++---- mjx/mujoco/mjx/_src/device.py | 2 +- mjx/mujoco/mjx/_src/device_test.py | 4 ++-- mjx/mujoco/mjx/_src/forward.py | 4 ++-- mjx/mujoco/mjx/_src/io.py | 4 ++-- mjx/mujoco/mjx/_src/io_test.py | 2 +- mjx/mujoco/mjx/_src/scan.py | 18 +++++++++--------- mjx/mujoco/mjx/_src/solver.py | 12 ++++++------ .../integration_test/collision_driver_test.py | 2 +- mjx/tutorial.ipynb | 2 +- 15 files changed, 52 insertions(+), 52 deletions(-) diff --git a/mjx/mujoco/mjx/_src/collision_convex.py b/mjx/mujoco/mjx/_src/collision_convex.py index 726ecade..193d999f 100644 --- a/mjx/mujoco/mjx/_src/collision_convex.py +++ b/mjx/mujoco/mjx/_src/collision_convex.py @@ -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() diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index 7fd905b4..1fabdf2f 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py index d501be5d..0f475aba 100644 --- a/mjx/mujoco/mjx/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/_src/collision_driver_test.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/collision_primitive.py b/mjx/mujoco/mjx/_src/collision_primitive.py index 7407919a..9040b22e 100644 --- a/mjx/mujoco/mjx/_src/collision_primitive.py +++ b/mjx/mujoco/mjx/_src/collision_primitive.py @@ -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)) ) diff --git a/mjx/mujoco/mjx/_src/collision_sdf.py b/mjx/mujoco/mjx/_src/collision_sdf.py index 999b3e8a..71a3dbf6 100644 --- a/mjx/mujoco/mjx/_src/collision_sdf.py +++ b/mjx/mujoco/mjx/_src/collision_sdf.py @@ -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)) ) diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 50aee9b3..f08e1988 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -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): diff --git a/mjx/mujoco/mjx/_src/device.py b/mjx/mujoco/mjx/_src/device.py index 0a8130e2..3ddb7a53 100644 --- a/mjx/mujoco/mjx/_src/device.py +++ b/mjx/mujoco/mjx/_src/device.py @@ -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: diff --git a/mjx/mujoco/mjx/_src/device_test.py b/mjx/mujoco/mjx/_src/device_test.py index 6126f707..530a4052 100644 --- a/mjx/mujoco/mjx/_src/device_test.py +++ b/mjx/mujoco/mjx/_src/device_test.py @@ -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): diff --git a/mjx/mujoco/mjx/_src/forward.py b/mjx/mujoco/mjx/_src/forward.py index f04eafe3..905e463a 100644 --- a/mjx/mujoco/mjx/_src/forward.py +++ b/mjx/mujoco/mjx/_src/forward.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 3e90b26e..9d0036b4 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 36806c2e..dde9c383 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -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) diff --git a/mjx/mujoco/mjx/_src/scan.py b/mjx/mujoco/mjx/_src/scan.py index 168637dc..a0cf3037 100644 --- a/mjx/mujoco/mjx/_src/scan.py +++ b/mjx/mujoco/mjx/_src/scan.py @@ -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)) diff --git a/mjx/mujoco/mjx/_src/solver.py b/mjx/mujoco/mjx/_src/solver.py index dae5a82b..540cd094 100644 --- a/mjx/mujoco/mjx/_src/solver.py +++ b/mjx/mujoco/mjx/_src/solver.py @@ -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) diff --git a/mjx/mujoco/mjx/integration_test/collision_driver_test.py b/mjx/mujoco/mjx/integration_test/collision_driver_test.py index a9891656..06116e67 100644 --- a/mjx/mujoco/mjx/integration_test/collision_driver_test.py +++ b/mjx/mujoco/mjx/integration_test/collision_driver_test.py @@ -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): diff --git a/mjx/tutorial.ipynb b/mjx/tutorial.ipynb index f57aee4d..a839c308 100644 --- a/mjx/tutorial.ipynb +++ b/mjx/tutorial.ipynb @@ -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",