import jax jax.config.update('jax_num_cpu_devices', 8) jax.devices() import numpy as np import jax.numpy as jnp arr = jnp.arange(32.0).reshape(4, 8) arr.devices() arr.sharding jax.debug.visualize_array_sharding(arr) from jax.sharding import PartitionSpec as P mesh = jax.make_mesh((2, 4), ('x', 'y')) sharding = jax.sharding.NamedSharding(mesh, P('x', 'y')) print(sharding) arr_sharded = jax.device_put(arr, sharding) print(arr_sharded) jax.debug.visualize_array_sharding(arr_sharded) @jax.jit def f_elementwise(x): return 2 * jnp.sin(x) + 1 result = f_elementwise(arr_sharded) print("shardings match:", result.sharding == arr_sharded.sharding) @jax.jit def f_contract(x): return x.sum(axis=0) result = f_contract(arr_sharded) jax.debug.visualize_array_sharding(result) print(result) some_array = np.arange(8) print(f"JAX-level type of some_array: {jax.typeof(some_array)}") @jax.jit def foo(x): print(f"JAX-level type of x during tracing: {jax.typeof(x)}") return x + x foo(some_array) from jax.sharding import AxisType mesh = jax.make_mesh((2, 4), ("X", "Y"), axis_types=(AxisType.Explicit, AxisType.Explicit)) replicated_array = np.arange(8).reshape(4, 2) sharded_array = jax.device_put(replicated_array, jax.NamedSharding(mesh, P("X", None))) print(f"replicated_array type: {jax.typeof(replicated_array)}") print(f"sharded_array type: {jax.typeof(sharded_array)}") arg0 = jax.device_put(np.arange(4).reshape(4, 1), jax.NamedSharding(mesh, P("X", None))) arg1 = jax.device_put(np.arange(8).reshape(1, 8), jax.NamedSharding(mesh, P(None, "Y"))) @jax.jit def add_arrays(x, y): ans = x + y print(f"x sharding: {jax.typeof(x)}") print(f"y sharding: {jax.typeof(y)}") print(f"ans sharding: {jax.typeof(ans)}") return ans with jax.set_mesh(mesh): add_arrays(arg0, arg1) mesh = jax.make_mesh((8,), ('x',)) f_elementwise_sharded = jax.shard_map( f_elementwise, mesh=mesh, in_specs=P('x'), out_specs=P('x')) arr = jnp.arange(32) f_elementwise_sharded(arr) x = jnp.arange(32) print(f"global shape: {x.shape=}") def f(x): print(f"device local shape: {x.shape=}") return x * 2 y = jax.shard_map(f, mesh=mesh, in_specs=P('x'), out_specs=P('x'))(x) def f(x): return jnp.sum(x, keepdims=True) jax.shard_map(f, mesh=mesh, in_specs=P('x'), out_specs=P('x'))(x) def f(x): sum_in_shard = x.sum() return jax.lax.psum(sum_in_shard, 'x') jax.shard_map(f, mesh=mesh, in_specs=P('x'), out_specs=P())(x) @jax.jit def layer(x, weights, bias): return jax.nn.sigmoid(x @ weights + bias) import numpy as np rng = np.random.default_rng(0) x = rng.normal(size=(32,)) weights = rng.normal(size=(32, 4)) bias = rng.normal(size=(4,)) layer(x, weights, bias) mesh = jax.make_mesh((8,), ('x',)) x_sharded = jax.device_put(x, jax.NamedSharding(mesh, P('x'))) weights_sharded = jax.device_put(weights, jax.NamedSharding(mesh, P())) layer(x_sharded, weights_sharded, bias) explicit_mesh = jax.make_mesh((8,), ('X',), axis_types=(AxisType.Explicit,)) x_sharded = jax.device_put(x, jax.NamedSharding(explicit_mesh, P('X'))) weights_sharded = jax.device_put(weights, jax.NamedSharding(explicit_mesh, P())) @jax.jit def layer_auto(x, weights, bias): print(f"x sharding: {jax.typeof(x)}") print(f"weights sharding: {jax.typeof(weights)}") print(f"bias sharding: {jax.typeof(bias)}") out = layer(x, weights, bias) print(f"out sharding: {jax.typeof(out)}") return out with jax.set_mesh(explicit_mesh): layer_auto(x_sharded, weights_sharded, bias) from functools import partial @jax.jit @partial(jax.shard_map, mesh=mesh, in_specs=(P('x'), P('x', None), P(None)), out_specs=P(None)) def layer_sharded(x, weights, bias): return jax.nn.sigmoid(jax.lax.psum(x @ weights, 'x') + bias) layer_sharded(x, weights, bias)