Spaces:
Sleeping
Sleeping
Fix: shim jax.device_put_replicated removed in JAX 0.10
#2
by arminfg - opened
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 |
|