jztree 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
jztree/__init__.py ADDED
@@ -0,0 +1,15 @@
1
+ from ._backend import load_backend as _load_backend
2
+
3
+ _jztree_cuda = _load_backend()
4
+
5
+ from . import jax_ext
6
+ from . import config
7
+ from . import stats
8
+ from . import data
9
+ from . import tools
10
+ from . import comm
11
+ from . import tree
12
+ from . import knn
13
+ from . import fof
14
+
15
+ del _load_backend
jztree/_backend.py ADDED
@@ -0,0 +1,92 @@
1
+ """Backend loader and CUDA compatibility checks for jztree."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import importlib.metadata
6
+ import os
7
+ import re
8
+ from typing import Optional
9
+
10
+
11
+ def _infer_jax_cuda_major() -> Optional[int]:
12
+ """Best-effort inference of JAX CUDA major version.
13
+
14
+ Returns None if JAX is missing or no reliable CUDA major can be determined.
15
+ """
16
+ try:
17
+ # Prefer explicit CUDA plugin package markers when available.
18
+ importlib.metadata.version("jax-cuda13-plugin")
19
+ return 13
20
+ except importlib.metadata.PackageNotFoundError:
21
+ pass
22
+
23
+ try:
24
+ importlib.metadata.version("jax-cuda12-plugin")
25
+ return 12
26
+ except importlib.metadata.PackageNotFoundError:
27
+ pass
28
+
29
+ try:
30
+ import jax
31
+ except Exception:
32
+ return None
33
+
34
+ try:
35
+ gpu_devices = jax.devices("gpu")
36
+ except Exception:
37
+ gpu_devices = []
38
+
39
+ if not gpu_devices:
40
+ return None
41
+
42
+ # Typical platform_version strings include tokens like "CUDA 13.0.1".
43
+ platform_version = str(getattr(gpu_devices[0], "platform_version", ""))
44
+ match = re.search(r"CUDA\s*([0-9]+)", platform_version, flags=re.IGNORECASE)
45
+ if match:
46
+ return int(match.group(1))
47
+
48
+ return None
49
+
50
+
51
+ def load_backend():
52
+ """Import jztree_cuda and validate CUDA major compatibility with JAX."""
53
+ skip_check = os.environ.get("JZTREE_SKIP_JAX_CUDA_CHECK", "").strip().lower()
54
+ skip_check = skip_check in {"1", "true", "yes", "on"}
55
+
56
+ try:
57
+ import jztree_cuda
58
+ except ImportError as exc:
59
+ raise ImportError(
60
+ "jztree backend package 'jztree_cuda' is not installed. "
61
+ "Install one backend wheel, e.g. 'jztree-cu12' or 'jztree-cu13'."
62
+ ) from exc
63
+
64
+ # Used by docs or lightweight environments to bypass backend compatibility checks.
65
+ if skip_check:
66
+ return jztree_cuda
67
+
68
+ backend_cuda_major = getattr(jztree_cuda, "CUDA_MAJOR", None)
69
+ if backend_cuda_major is None:
70
+ raise ImportError(
71
+ "Installed jztree backend does not expose CUDA_MAJOR. "
72
+ "Please reinstall a supported backend wheel (jztree-cu12 or jztree-cu13)."
73
+ )
74
+
75
+ try:
76
+ backend_cuda_major = int(backend_cuda_major)
77
+ except Exception as exc:
78
+ raise ImportError(
79
+ f"Installed jztree backend has invalid CUDA_MAJOR={backend_cuda_major!r}. "
80
+ "Please reinstall a supported backend wheel (jztree-cu12 or jztree-cu13)."
81
+ ) from exc
82
+
83
+ if not skip_check:
84
+ jax_cuda_major = _infer_jax_cuda_major()
85
+ if jax_cuda_major is not None and jax_cuda_major != backend_cuda_major:
86
+ raise ImportError(
87
+ "Installed jztree backend is CUDA "
88
+ f"{backend_cuda_major}, but JAX appears to use CUDA {jax_cuda_major}. "
89
+ f"Please install jztree-cu{jax_cuda_major}."
90
+ )
91
+
92
+ return jztree_cuda
jztree/_version.py ADDED
@@ -0,0 +1,24 @@
1
+ # file generated by vcs-versioning
2
+ # don't change, don't track in version control
3
+ from __future__ import annotations
4
+
5
+ __all__ = [
6
+ "__version__",
7
+ "__version_tuple__",
8
+ "version",
9
+ "version_tuple",
10
+ "__commit_id__",
11
+ "commit_id",
12
+ ]
13
+
14
+ version: str
15
+ __version__: str
16
+ __version_tuple__: tuple[int | str, ...]
17
+ version_tuple: tuple[int | str, ...]
18
+ commit_id: str | None
19
+ __commit_id__: str | None
20
+
21
+ __version__ = version = '0.4.dev263+g79edea172.d20260406'
22
+ __version_tuple__ = version_tuple = (0, 4, 'dev263', 'g79edea172.d20260406')
23
+
24
+ __commit_id__ = commit_id = 'g79edea172'
jztree/comm.py ADDED
@@ -0,0 +1,567 @@
1
+ from dataclasses import dataclass
2
+ import numpy as np
3
+ import os
4
+ import jax
5
+ import jax.numpy as jnp
6
+ from typing import Tuple, Any, TypeAlias
7
+
8
+ from .jax_ext import get_rank_info, tree_map_by_len, raise_if
9
+ from .tools import cumsum_starting_with_zero, inverse_of_splits, ragged_transpose
10
+ from .jax_ext import empty_like, invalidate, pytree_len, leading_len
11
+
12
+ # Currently jax doesn't have a typehint for pytrees. We simply define one ourselves for clarity
13
+ Pytree: TypeAlias = Any
14
+
15
+ # ------------------------------------------------------------------------------------------------ #
16
+ # Distributed Initialization Helpers #
17
+ # ------------------------------------------------------------------------------------------------ #
18
+
19
+ def _env_int(name: str, default: int = 0) -> int:
20
+ v = os.environ.get(name)
21
+ if v is None:
22
+ return default
23
+ try:
24
+ return int(v)
25
+ except ValueError:
26
+ return default
27
+
28
+ def should_init_jax_distributed() -> bool:
29
+ # --- Slurm ---
30
+ # SLURM_NTASKS is total tasks; if >1 we’re distributed.
31
+ if _env_int("SLURM_NTASKS", 1) > 1:
32
+ return True
33
+
34
+ # --- Open MPI ---
35
+ if _env_int("OMPI_COMM_WORLD_SIZE", 1) > 1:
36
+ return True
37
+
38
+ # --- MPICH / PMI-based launchers (incl. some Slurm/MPI setups) ---
39
+ if _env_int("PMI_SIZE", 1) > 1 or _env_int("PMIX_SIZE", 1) > 1:
40
+ return True
41
+
42
+ # --- Generic “world size” used by some launchers ---
43
+ if _env_int("WORLD_SIZE", 1) > 1:
44
+ return True
45
+
46
+ # --- “Explicit JAX distributed config present” heuristic ---
47
+ # If you set these yourself in your job wrapper, treat it as distributed.
48
+ if any(k in os.environ for k in ("JAX_PROCESS_COUNT", "JAX_PROCESS_INDEX", "JAX_COORDINATOR_ADDRESS")):
49
+ # Only treat it as distributed if it’s actually >1 (when provided).
50
+ pc = _env_int("JAX_PROCESS_COUNT", 1)
51
+ return pc > 1
52
+
53
+ return False
54
+
55
+ # ------------------------------------------------------------------------------------------------ #
56
+ # Tiny Helper Functions #
57
+ # ------------------------------------------------------------------------------------------------ #
58
+
59
+
60
+ # ------------------------------------------------------------------------------------------------ #
61
+ # Packing Helpers for more efficient comm #
62
+ # ------------------------------------------------------------------------------------------------ #
63
+
64
+ @dataclass
65
+ class PackingSpec:
66
+ treedef: Any
67
+ shapes: Tuple[Tuple[int, ...], ...]
68
+ dtypes: Tuple[jnp.dtype, ...]
69
+ offsets: np.ndarray
70
+ keep_mask: Tuple[bool, ...] # True if leaf was packed (shape[0] == N)
71
+
72
+ def _pack_pytree(x: Pytree, N: int) -> Tuple[jax.Array, PackingSpec]:
73
+ """Packs only leaves with leading size N (shape[0] == N). Other leaves are skipped."""
74
+ leaves, treedef = jax.tree_util.tree_flatten(x)
75
+
76
+ def keep(l):
77
+ return (
78
+ hasattr(l, "shape") and hasattr(l, "dtype")
79
+ and l.shape is not None and len(l.shape) >= 1
80
+ and int(l.shape[0]) == int(N)
81
+ )
82
+
83
+ keep_mask = tuple(keep(l) for l in leaves)
84
+ data_leaves = [l for l, m in zip(leaves, keep_mask) if m]
85
+
86
+ assert len(data_leaves) > 0
87
+
88
+ base_shape = (int(N),)
89
+ for l in data_leaves:
90
+ assert l.shape[:1] == base_shape, "packed leaves must align in leading dimension"
91
+
92
+ dtypes = tuple(l.dtype for l in data_leaves)
93
+ shapes = tuple(l.shape for l in data_leaves)
94
+
95
+ arrs = [l.reshape((N, -1)).view(jnp.uint8) for l in data_leaves]
96
+ num = [a.shape[-1] for a in arrs]
97
+ with jax.enable_x64():
98
+ offsets = np.pad(np.cumsum(num), (1, 0)).astype(np.int64)
99
+
100
+ spec = PackingSpec(treedef, shapes, dtypes, offsets, keep_mask)
101
+ return jnp.concatenate(arrs, axis=-1), spec
102
+
103
+ def _unpack_pytree(x: jax.Array, p: PackingSpec, template: Pytree) -> Pytree:
104
+ """
105
+ Unpacks packed leaves and merges into `template`.
106
+ Leaves that were skipped during packing are taken from `template` unchanged.
107
+ """
108
+ # Rebuild packed leaves
109
+ packed_leaves = []
110
+ for i in range(len(p.shapes)):
111
+ sl = x[..., p.offsets[i]:p.offsets[i+1]]
112
+ packed_leaves.append(sl.view(p.dtypes[i]).reshape(p.shapes[i]))
113
+
114
+ tmpl_leaves, tmpl_def = jax.tree_util.tree_flatten(template)
115
+ if tmpl_def != p.treedef:
116
+ raise ValueError("Template treedef mismatch.")
117
+
118
+ # Merge packed leaves into template leaves
119
+ out_leaves = []
120
+ pi = 0
121
+ for m, tleaf in zip(p.keep_mask, tmpl_leaves):
122
+ if m:
123
+ out_leaves.append(packed_leaves[pi])
124
+ pi += 1
125
+ else:
126
+ out_leaves.append(tleaf)
127
+
128
+ return jax.tree_util.tree_unflatten(p.treedef, out_leaves)
129
+ # ------------------------------------------------------------------------------------------------ #
130
+ # Simple Communication Directives #
131
+ # ------------------------------------------------------------------------------------------------ #
132
+
133
+ def global_splits(n, axis_name=None):
134
+ if axis_name is None:
135
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
136
+ alln = jax.lax.all_gather(n, axis_name)
137
+ return jnp.pad(jnp.cumsum(alln), (1,0), constant_values=0)
138
+
139
+ def send_to_right(x, axis_name, invalid_float=jnp.nan, invalid_int=0):
140
+ rank = jax.lax.axis_index(axis_name)
141
+ ndev = jax.lax.axis_size(axis_name)
142
+
143
+ xin = jax.lax.ppermute(x, axis_name, [(i, i+1) for i in range(0,ndev-1)])
144
+ xin = invalidate(xin, rank == 0, invalid_float, invalid_int)
145
+
146
+ return xin
147
+
148
+ def send_to_left(x, axis_name, invalid_float=jnp.nan, invalid_int=0):
149
+ rank = jax.lax.axis_index(axis_name)
150
+ ndev = jax.lax.axis_size(axis_name)
151
+
152
+ xin = jax.lax.ppermute(x, axis_name, [(i, i-1) for i in range(1,ndev)])
153
+ xin = invalidate(xin, rank == ndev-1, invalid_float, invalid_int)
154
+
155
+ return xin
156
+
157
+ def get_pos(x):
158
+ if isinstance(x, jax.Array):
159
+ return x
160
+ else: # assume x is a pytree with .pos attribute
161
+ return x.pos
162
+
163
+ def shift_particles_left(x, nsend, max_send, npart):
164
+ """Sends the first nsend elements of every leaf with size == largest leaf size"""
165
+ rank, ndev, axis_name = get_rank_info()
166
+
167
+ # Validate that send buffer is large enough
168
+ npart = npart + raise_if(nsend >= max_send,
169
+ "Cannot fit {nsend} particles into buffer of size {max_send}!",
170
+ nsend=nsend, max_send=max_send
171
+ )
172
+
173
+ # Validate that particle array has enough free space
174
+ nget = send_to_left(nsend, axis_name)
175
+ npart = npart + raise_if(npart + nget - nsend >= pytree_len(x),
176
+ "Cannot shift particles: have={nhave}, get={nget}, send={nsend}, max={nmax}.",
177
+ nhave=npart, nget=nget, nsend=nsend, nmax=pytree_len(x)
178
+ )
179
+
180
+ # Send the particles
181
+ size = pytree_len(x)
182
+
183
+ x_get = send_to_left(tree_map_by_len(lambda v: v[0:max_send], x, size), axis_name, invalid_float=jnp.nan)
184
+
185
+ # Delete the particles that were send
186
+ iar = jnp.arange(pytree_len(x))
187
+ x = invalidate(x, iar < nsend)
188
+ x = tree_map_by_len(lambda v: jnp.roll(v, -nsend, axis=0), x, size)
189
+
190
+ # Insert the received particles
191
+ idx = jnp.arange(max_send)
192
+ idx = jnp.where(idx < nget, npart - nsend + idx, pytree_len(x)) # discard indices beyond nadd
193
+ def insert(u, v):
194
+ if leading_len(u) != size:
195
+ return u
196
+ else:
197
+ return u.at[idx].set(v)
198
+ x = jax.tree.map(insert, x, x_get)
199
+
200
+ return x, npart + nget - nsend
201
+
202
+ # ------------------------------------------------------------------------------------------------ #
203
+ # All To All communication #
204
+ # ------------------------------------------------------------------------------------------------ #
205
+
206
+ def ragged_all_to_all_through_buf(operand, output, input_offsets, send_sizes, output_offsets, recv_sizes, *, axis_name, buf_size=1024):
207
+ """Does the same as jax.lax.ragged_all_to_all, but works on CPU and needs a buffer"""
208
+ ndev = jax.lax.axis_size(axis_name)
209
+ output_offsets = jax.lax.all_to_all(output_offsets, axis_name, 0, 0, tiled=True)
210
+
211
+ xp, xspec = _pack_pytree(operand)
212
+
213
+ def comm(icom, output):
214
+ igpu, ioff = jnp.indices((ndev, buf_size))
215
+ inoff = input_offsets[igpu] + icom*buf_size + ioff
216
+ xbuf = xp[inoff]
217
+
218
+ xrecv = jax.lax.all_to_all(xbuf, axis_name, 0, 0, tiled=True)
219
+
220
+ valid = icom*buf_size + ioff < recv_sizes[igpu]
221
+ outoff = output_offsets[igpu] + icom*buf_size + ioff
222
+ outoff = jnp.where(valid, outoff, output.size)
223
+ output = output.at[outoff].set(xrecv)
224
+
225
+ return output
226
+
227
+ max_size = jax.lax.pmax(jnp.max(send_sizes), axis_name)
228
+ ncomm = (max_size + buf_size - 1) // buf_size
229
+
230
+ op, ospec = _pack_pytree(output, pytree_len(output))
231
+ op = jax.lax.fori_loop(0, ncomm, comm, op)
232
+ return _unpack_pytree(op, ospec)
233
+
234
+ def all_to_all_with_splits(x, ispl, output=None, axis_name=None, verify=True, err_hint="", copy_self=True, pack_pytree=False):
235
+ """all_to_all communication with data-dependent communication volume
236
+
237
+ We send to rank i: x[ispl[i]:ispl[i+1]]
238
+ all received values will be inserted continguously into output (starting at 0).
239
+
240
+ If x is a pytree, communication will only be applied to those leaves whoes length corresponds
241
+ to the length of the largest leaf
242
+
243
+ x: jax.Array or pytree. If it is a pytree the communication will be applied over the leading
244
+ dimensions of all leaves (undefined behaviour if some leaves have different lengths)
245
+ output: jax.Array or pytree. If x is a pytree output needs to be of identical structure.
246
+ If not provided, we use a copy of x filled with jnp.nan (or 0 for integers)
247
+ verify: If True, throws an error if output buffer is too small. Otherwise out-of-range values
248
+ will simply be discarded.
249
+ err_hint: If given, add a hint to the potential error message, indicating how to fix it
250
+ copy_self: Extract self-send data and copy it directly (surprisingly this is faster)
251
+ """
252
+ if output is None:
253
+ output = empty_like(x)
254
+ if axis_name is None:
255
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
256
+
257
+ out_size = pytree_len(output)
258
+
259
+ input_offsets = ispl[:-1]
260
+ send_sizes = ispl[1:] - ispl[:-1]
261
+ recv_sizes = jax.lax.all_to_all(send_sizes, axis_name, 0, 0, tiled=True)
262
+ dev_spl = cumsum_starting_with_zero(recv_sizes)
263
+ output_offsets = dev_spl[:-1]
264
+
265
+ if verify:
266
+ need = jnp.sum(recv_sizes)
267
+ recv_sizes = recv_sizes + raise_if(need > out_size,
268
+ "The receiving buffer is too small, need={need}, have={have}" + err_hint,
269
+ need=need, have=out_size
270
+ )
271
+
272
+ if copy_self:
273
+ # avoid communication for self i/o
274
+ rank = jax.lax.axis_index(axis_name)
275
+ iout = jnp.arange(out_size) #+ output_offsets[rank]
276
+ iin = jnp.arange(out_size) + input_offsets[rank] - output_offsets[rank]
277
+ mask = (iout >= output_offsets[rank]) & (iout < output_offsets[rank] + send_sizes[rank])
278
+
279
+ send_sizes = send_sizes.at[rank].set(0)
280
+ recv_sizes = recv_sizes.at[rank].set(0)
281
+
282
+ def copy(xi, outi):
283
+ if leading_len(xi) != leading_len(mask):
284
+ return outi
285
+ mask_rs = jnp.reshape(mask, (len(mask),) + (1,)*(xi.ndim -1))
286
+ return jnp.where(mask_rs, xi[iin], outi)
287
+
288
+ output = jax.tree.map(copy, x, output)
289
+
290
+ # funnily jax.lax.ragged_all_to_all wants to know the output_offsets on the
291
+ # sending GPU rather than the receiving one... So we need to communicate the offsets
292
+ output_offsets = jax.lax.all_to_all(output_offsets, axis_name, 0, 0, tiled=True)
293
+
294
+ def comm(xi, outi):
295
+ if leading_len(outi) != out_size:
296
+ return outi
297
+ return jax.lax.ragged_all_to_all(
298
+ xi, outi, input_offsets, send_sizes, output_offsets, recv_sizes, axis_name=axis_name
299
+ )
300
+ # uncomment this for CPU
301
+ # return ragged_all_to_all_through_buf(
302
+ # xi, outi, input_offsets, send_sizes, output_offsets, recv_sizes, axis_name=axis_name
303
+ # )
304
+
305
+ if pack_pytree:
306
+ xp, xspec = _pack_pytree(x, pytree_len(x))
307
+ op, ospec = _pack_pytree(output, pytree_len(output))
308
+ op = comm(xp, op)
309
+ return _unpack_pytree(op, ospec, x), dev_spl
310
+ else:
311
+ return jax.tree.map(comm, x, output), dev_spl
312
+
313
+ def all_to_all_along_axis(data, nij, axis=1, err_hint="", copy_self=True, pack_pytree=False):
314
+ axis_names = jax.sharding.get_abstract_mesh().axis_names
315
+ shape = jax.sharding.get_abstract_mesh().axis_sizes
316
+ ndim = len(shape)
317
+
318
+ if shape[axis] == 1:
319
+ return data, nij
320
+
321
+ assert nij.ndim == ndim
322
+
323
+ axis_name = axis_names[axis]
324
+
325
+ # transposition that moves axis to the beginning
326
+ transpose = (axis,) + tuple(i for i in range(ndim) if i !=axis)
327
+ # transposition that moves it back to where it came from:
328
+ inv_transpose = tuple((i+(i<axis) if i != axis else 0) for i in range(ndim))
329
+
330
+ def get_splits(n):
331
+ ninner = jnp.sum(n, axis=range(1,ndim))
332
+ return jnp.pad(jnp.cumsum(ninner), (1,0), constant_values=0)
333
+
334
+ # transpose communication axis to beginning
335
+ data_ji, nji = ragged_transpose(data, nij, axes=transpose)
336
+
337
+ # communicate
338
+ data_ji, dspl = all_to_all_with_splits(
339
+ data_ji, get_splits(nji), axis_name=axis_name, err_hint=err_hint, copy_self=copy_self,
340
+ pack_pytree=pack_pytree
341
+ )
342
+ nji = jax.lax.all_to_all(nij, axis_name, axis, 0)
343
+
344
+ # transpose the axis back where it belongs
345
+ data_ij, nij = ragged_transpose(data_ji, nji, axes=inv_transpose)
346
+
347
+ return data_ij, nij
348
+
349
+ def nested_all_to_all_with_splits(data, ispl, **kwargs):
350
+ shape = jax.sharding.get_abstract_mesh().axis_sizes
351
+ ndim = len(shape)
352
+
353
+ nij = (ispl[1:] - ispl[:-1]).reshape(shape)
354
+
355
+ for i in reversed(range(0, ndim)):
356
+ data, nij = all_to_all_along_axis(data, nij, axis=i, **kwargs)
357
+
358
+ return data, jnp.pad(jnp.cumsum(nij.flatten()), (1,0), constant_values=0)
359
+
360
+ def dynamic_all_gather(x, nsend, output=None, axis_name=None, verify=True):
361
+ """An all-gather where each task may send different amounts.
362
+
363
+ returns output, dev_spl -- where output[dev_spl[i]:dev_spl[i+1]] contains rank i's input
364
+ """
365
+ if output is None:
366
+ output = empty_like(x)
367
+ if axis_name is None:
368
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
369
+
370
+ rank = jax.lax.axis_index(axis_name)
371
+ ndev = jax.lax.axis_size(axis_name)
372
+
373
+ nrecv = jax.lax.all_gather(nsend, axis_name)
374
+ dev_spl = cumsum_starting_with_zero(nrecv)
375
+
376
+ if verify:
377
+ out_size = pytree_len(output)
378
+ dev_spl = dev_spl + raise_if(dev_spl[-1] >= out_size,
379
+ "The receiveing buffer (size: {out_size}) is to small (need: {need})",
380
+ out_size=out_size, need=dev_spl[-1],
381
+ )
382
+
383
+ input_off = jnp.zeros(ndev, jnp.int32)
384
+ nsend = jnp.full(ndev, nsend, dtype=jnp.int32)
385
+ output_off = jnp.full(ndev, dev_spl[rank])
386
+
387
+ def comm(xi, outi):
388
+ return jax.lax.ragged_all_to_all(
389
+ xi, outi, input_off, nsend, output_off, nrecv, axis_name=axis_name
390
+ )
391
+
392
+ return jax.tree.map(comm, x, output), dev_spl
393
+
394
+ def arange_for_comm(irank: jax.Array, x: jax.Array,
395
+ num: jax.Array | int |None = None, axis_name=None):
396
+ if axis_name is None:
397
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
398
+ rank = jax.lax.axis_index(axis_name)
399
+ ndev = jax.lax.axis_size(axis_name)
400
+
401
+ if num is not None:
402
+ irank = jnp.where(jnp.arange(len(irank), dtype=irank.dtype) < num, irank, ndev)
403
+ isort = jnp.argsort(irank)
404
+ dev_spl = jnp.searchsorted(irank[isort], jnp.arange(ndev+1, dtype=irank.dtype), side="left")
405
+
406
+ xsort = tree_map_by_len(lambda d: d[isort], x, pytree_len(x))
407
+
408
+ return xsort, dev_spl, isort
409
+
410
+ def all_to_all_with_irank(
411
+ irank: jax.Array,
412
+ x: jax.Array | Pytree,
413
+ output: jax.Array | Pytree | None = None,
414
+ num: jax.Array | int | None = None,
415
+ axis_name: str = None,
416
+ verify: bool = True,
417
+ err_hint: str = "",
418
+ copy_self: bool = True,
419
+ pack_pytree: bool = False,
420
+ get_inverse: bool = False
421
+ ):
422
+ """Communicate by indicating the rank of the receiving device
423
+
424
+ To understand most arguments, see documentation of all_to_all_with_splits
425
+ """
426
+ if axis_name is None:
427
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
428
+ xsort, dev_spl, isort = arange_for_comm(irank, x, num=num, axis_name=axis_name)
429
+ x, dev_spl = all_to_all_with_splits(xsort, dev_spl, output, axis_name, verify=verify, err_hint=err_hint, copy_self=copy_self, pack_pytree=pack_pytree)
430
+ if get_inverse:
431
+ invsort = jnp.zeros_like(isort).at[isort].set(jnp.arange(len(isort), dtype=isort.dtype))
432
+ return x, dev_spl, invsort
433
+ else:
434
+ return x, dev_spl
435
+
436
+ def all_to_all_request(
437
+ irank: jax.Array,
438
+ indices: jax.Array,
439
+ x: jax.Array | Pytree,
440
+ output: jax.Array | Pytree | None = None,
441
+ num: jax.Array | int | None = None,
442
+ axis_name: str = None,
443
+ verify: bool = True,
444
+ err_hint: str = "",
445
+ copy_self: bool = True,
446
+ pack_pytree: bool = False
447
+ ) -> jax.Array:
448
+ if axis_name is None:
449
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
450
+ # First inform the task with the data which indices we need
451
+ indices_sort, dev_spl, isort = arange_for_comm(irank, indices, num=num, axis_name=axis_name)
452
+ indices, dev_spl = all_to_all_with_splits(
453
+ indices_sort, dev_spl, output=None, axis_name=axis_name, verify=verify, err_hint=err_hint,
454
+ copy_self=copy_self, pack_pytree=pack_pytree
455
+ )
456
+ # Then send back the data at those locations
457
+ xsort, dev_spl = all_to_all_with_splits(
458
+ x[indices], dev_spl, output, axis_name=axis_name, verify=verify, err_hint=err_hint,
459
+ copy_self=copy_self, pack_pytree=pack_pytree
460
+ )
461
+ # rearange to the original order
462
+ invsort = jnp.zeros_like(isort).at[isort].set(jnp.arange(len(isort), dtype=isort.dtype))
463
+ return jax.tree.map(lambda xi: xi[invsort], xsort)
464
+
465
+ def all_to_all_request_children(
466
+ dev_spl: jax.Array,
467
+ indices: jax.Array,
468
+ spl: jax.Array,
469
+ data: jax.Array | Pytree,
470
+ output: jax.Array | Pytree | None = None,
471
+ axis_name: str = None,
472
+ verify: bool = True,
473
+ err_hint_parent: str = "",
474
+ err_hint_child: str = "",
475
+ copy_self: bool = True,
476
+ pack_pytree: bool = False
477
+ ):
478
+ if output is None:
479
+ output = empty_like(data)
480
+ if axis_name is None:
481
+ axis_name = jax.sharding.get_abstract_mesh().axis_names
482
+
483
+ size = pytree_len(output)
484
+
485
+ # First inform the task with the data which indices we need
486
+ indices, dev_spl = all_to_all_with_splits(
487
+ indices, dev_spl, output=None, axis_name=axis_name, verify=verify, err_hint=err_hint_parent,
488
+ copy_self=copy_self, pack_pytree=pack_pytree
489
+ )
490
+
491
+ # fill a continous buffer with the requested data
492
+ node_sizes = (spl[1:] - spl[:-1])[indices]
493
+ out_node_spl = cumsum_starting_with_zero(node_sizes)
494
+ child_inode = inverse_of_splits(out_node_spl, size)
495
+ child_inode_offset = jnp.arange(size) - out_node_spl[child_inode]
496
+ child_id = spl[indices[child_inode]] + child_inode_offset
497
+ child_data = tree_map_by_len(lambda xi: xi[child_id], data, pytree_len(data))
498
+ child_dev_spl = out_node_spl[dev_spl]
499
+
500
+ # send back the node_sizes so the receiver knows where each node starts
501
+ node_sizes, node_dev_spl = all_to_all_with_splits(
502
+ node_sizes, dev_spl, axis_name=axis_name, verify=verify,
503
+ err_hint="\nThis should not fail...", # (since it must be identical to input size)
504
+ pack_pytree=pack_pytree,
505
+ )
506
+ node_spl = cumsum_starting_with_zero(node_sizes)
507
+
508
+ # Now send the child data
509
+ xchild, child_dev_spl = all_to_all_with_splits(
510
+ child_data, child_dev_spl, output, axis_name=axis_name, verify=verify, err_hint=err_hint_child,
511
+ copy_self=copy_self, pack_pytree=pack_pytree
512
+ )
513
+
514
+ return xchild, node_spl, child_dev_spl
515
+
516
+ # ------------------------------------------------------------------------------------------------ #
517
+ # All to all with permute #
518
+ # ------------------------------------------------------------------------------------------------ #
519
+
520
+ def update_range(x, update, i1, i2):
521
+ idx = i1 + jnp.arange(len(update))
522
+ return x.at[jnp.where(idx < i2, idx, len(x))].set(update)
523
+
524
+ def gather_num(x, size, at):
525
+ return x[at + jnp.arange(size)]
526
+
527
+ def permute_offset(offset, num):
528
+ return [(i, (i + offset) % num) for i in range(num)]
529
+
530
+ def all_to_all_with_permute(x, ispl, buffer_bytes=8*1024**2, axis_name=None, verify=True):
531
+ rank, ndev, axis_name = get_rank_info(axis_name)
532
+
533
+ x, spec = _pack_pytree(x, pytree_len(x))
534
+ bsize = max(buffer_bytes // (x[0].size*x.itemsize), 1)
535
+
536
+ output = jnp.copy(x)
537
+
538
+ nsend = ispl[1:] - ispl[:-1]
539
+ nrecv = jax.lax.all_to_all(nsend, axis_name, split_axis=0, concat_axis=0, tiled=True)
540
+ ispl_recv = jnp.pad(jnp.cumsum(nrecv), (1,0))
541
+
542
+ if verify:
543
+ out_size = pytree_len(output)
544
+ ispl_recv = ispl_recv + raise_if(ispl_recv[-1] >= out_size,
545
+ "The receiving buffer (size: {out_size}) is to small (need: {need})",
546
+ out_size=out_size, need=ispl_recv[-1],
547
+ )
548
+
549
+ nsteps = (jnp.roll(nsend, -rank) + bsize - 1) // bsize
550
+ nsteps = jax.lax.pmax(nsteps, axis_name)
551
+
552
+ # jax.debug.print("s1 {}, s2 {}, nsteps {}", nsend, bsize, nsteps)
553
+
554
+ def handle_offset(output, offset, nsteps):
555
+ ito, ifrom = (rank + offset) % ndev, (rank - offset) % ndev
556
+ def step(i, output):
557
+ data_send = gather_num(x, bsize, at=ispl[ito]+bsize*i)
558
+ data_recv = jax.lax.ppermute(data_send, axis_name, permute_offset(offset, ndev))
559
+ return update_range(output, data_recv, ispl_recv[ifrom] + bsize*i, ispl_recv[ifrom+1])
560
+ return jax.lax.fori_loop(0, nsteps, step, output)
561
+
562
+ output = update_range(output, gather_num(x, len(x), at=ispl[rank]), ispl_recv[rank], ispl_recv[rank+1])
563
+
564
+ for offset in range(1, ndev):
565
+ output = handle_offset(output, offset, nsteps[offset])
566
+
567
+ return _unpack_pytree(output, spec, x), ispl_recv