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 +15 -0
- jztree/_backend.py +92 -0
- jztree/_version.py +24 -0
- jztree/comm.py +567 -0
- jztree/config.py +65 -0
- jztree/data.py +686 -0
- jztree/fof.py +916 -0
- jztree/jax_ext.py +361 -0
- jztree/knn.py +356 -0
- jztree/stats.py +174 -0
- jztree/tools.py +277 -0
- jztree/tree.py +995 -0
- jztree-1.0.0.dist-info/METADATA +18 -0
- jztree-1.0.0.dist-info/RECORD +16 -0
- jztree-1.0.0.dist-info/WHEEL +5 -0
- jztree-1.0.0.dist-info/top_level.txt +1 -0
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
|