Skip to content

Commit a22fe93

Browse files
danielsuocopybara-github
authored andcommitted
[pmap] In-line definitions of jax.device_put_sharded and jax.device_put_replicated.
Both `jax.device_put_sharded` and `jax.device_put_replicated` were deprecated in JAX v0.8.1 in November 2025. We in-line their definitions using public JAX APIs taking the `jax_pmap_shmap_merge=True` branch, which was made the default in JAX v0.8.0 in October 2025. Please see the below for more information: - JAX CHANGELOG: https://docs.jax.dev/en/latest/changelog.html - Migrating from `jax.pmap`: https://docs.jax.dev/en/latest/migrate_pmap.html PiperOrigin-RevId: 888582443
1 parent 98963e2 commit a22fe93

1 file changed

Lines changed: 14 additions & 2 deletions

File tree

‎clrs/_src/baselines.py‎

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -125,15 +125,27 @@ def _maybe_put_replicated(tree):
125125
if jax.local_device_count() == 1:
126126
return jax.device_put(tree)
127127
else:
128-
return jax.device_put_replicated(tree, jax.local_devices())
128+
devices = jax.local_devices()
129+
mesh = jax.sharding.Mesh(np.array(devices), ('_device_put_sharded',))
130+
sharding = jax.NamedSharding(mesh, jax.P('_device_put_sharded'))
131+
132+
def _replicate(x):
133+
if isinstance(x, jax.Array):
134+
return jax.device_put(jnp.stack([x] * len(devices)), sharding)
135+
return jax.device_put(np.stack([x] * len(devices)), sharding)
136+
137+
return jax.tree_util.tree_map(_replicate, tree)
129138

130139

131140
def _maybe_pmap_rng_key(rng_key: _Array):
132141
n_devices = jax.local_device_count()
133142
if n_devices == 1:
134143
return rng_key
144+
devices = jax.local_devices()
135145
pmap_rng_keys = jax.random.split(rng_key, n_devices)
136-
return jax.device_put_sharded(list(pmap_rng_keys), jax.local_devices())
146+
mesh = jax.sharding.Mesh(np.array(devices), ('_device_put_sharded',))
147+
sharding = jax.NamedSharding(mesh, jax.P('_device_put_sharded'))
148+
return jax.device_put(jnp.stack(list(pmap_rng_keys)), sharding)
137149

138150

139151
class BaselineModel(model.Model):

0 commit comments

Comments
 (0)