Fix: shim jax.device_put_replicated removed in JAX 0.10

#2
Files changed (1) hide show
  1. app.py +48 -1
app.py CHANGED
@@ -53,6 +53,50 @@ def _upload(path: Path) -> None:
53
  log(f"upload failed ({type(e).__name__}): {e}")
54
 
55
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
56
  def train_worker() -> None:
57
  try:
58
  STATE["status"] = "importing"
@@ -60,7 +104,10 @@ def train_worker() -> None:
60
  import functools
61
  import jax
62
 
63
- log(f"jax devices: {jax.devices()}")
 
 
 
64
  if not any(d.platform == "gpu" for d in jax.devices()):
65
  log("WARNING: no GPU visible to JAX -- this will be very slow")
66
 
 
53
  log(f"upload failed ({type(e).__name__}): {e}")
54
 
55
 
56
+ def install_jax_pmap_shims() -> None:
57
+ """Re-add jax.device_put_replicated / device_put_sharded if this JAX removed them.
58
+
59
+ JAX deprecated both in 0.8.1 and removed them in 0.10.0 (April 2026), but the
60
+ current brax *release* (0.14.2, which playground 0.2.0 requires) still calls
61
+ device_put_replicated; only brax main has the sharding-based replacement.
62
+ Shimming the two functions is a smaller intervention than pinning JAX down,
63
+ which would risk the MJX/Warp stack that already compiles cleanly here.
64
+ """
65
+ import jax
66
+ import jax.numpy as jnp
67
+ import numpy as np
68
+
69
+ def _sharding(devices):
70
+ mesh = jax.sharding.Mesh(np.array(list(devices)), axis_names=("i",))
71
+ return jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec("i"))
72
+
73
+ def device_put_replicated(x, devices):
74
+ sharding, n = _sharding(devices), len(devices)
75
+
76
+ def rep(leaf):
77
+ stack = jnp.stack if isinstance(leaf, jax.Array) else np.stack
78
+ return jax.device_put(stack([leaf] * n), sharding)
79
+
80
+ return jax.tree_util.tree_map(rep, x)
81
+
82
+ def device_put_sharded(shards, devices):
83
+ sharding = _sharding(devices)
84
+
85
+ def put(*leaves):
86
+ stack = jnp.stack if isinstance(leaves[0], jax.Array) else np.stack
87
+ return jax.device_put(stack(list(leaves)), sharding)
88
+
89
+ return jax.tree_util.tree_map(put, *shards)
90
+
91
+ for name, fn in (("device_put_replicated", device_put_replicated),
92
+ ("device_put_sharded", device_put_sharded)):
93
+ try:
94
+ getattr(jax, name)
95
+ except AttributeError:
96
+ setattr(jax, name, fn)
97
+ log(f"shimmed jax.{name} (removed in this JAX version)")
98
+
99
+
100
  def train_worker() -> None:
101
  try:
102
  STATE["status"] = "importing"
 
104
  import functools
105
  import jax
106
 
107
+ log(f"jax {jax.__version__} devices: {jax.devices()}")
108
+ install_jax_pmap_shims()
109
+ import brax
110
+ log(f"brax {brax.__version__}")
111
  if not any(d.platform == "gpu" for d in jax.devices()):
112
  log("WARNING: no GPU visible to JAX -- this will be very slow")
113