File size: 12,963 Bytes
5ccb4fd
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
# Copyright 2022 DeepMind Technologies Limited. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""K-FAC utilities for multi-device execution."""
import functools
from typing import Any, Callable, Sequence

import jax
from jax import lax
import jax.numpy as jnp
from kfac_jax._src.utils import types

try:
  # JAX v0.10.0 or newer
  from jax.extend.core import unsafe_get_axis_names_DO_NOT_USE  # pylint: disable=g-import-not-at-top
except ImportError:
  # JAX v0.9.2 or older
  from jax.core import unsafe_get_axis_names_DO_NOT_USE  # pylint: disable=g-import-not-at-top

jax_version = (
    jax.__version_info__ if hasattr(jax, "__version_info__")
    else tuple(map(int, jax.__version__.split("."))))


Array = types.Array
Numeric = types.Numeric
PRNGKey = types.PRNGKey
TArrayTree = types.TArrayTree


def _axis_name_tuple(axis_name):
  if axis_name is None:
    return ()
  if isinstance(axis_name, tuple):
    return axis_name
  return (axis_name,)


def in_pmap(axis_name: str | tuple[str, ...] | None) -> bool:
  """Returns whether we are in a pmap with the given axis name."""

  if axis_name is None:
    return False

  axis_names = unsafe_get_axis_names_DO_NOT_USE()
  requested = _axis_name_tuple(axis_name)

  if all(name in axis_names for name in requested):
    return True

  if len(axis_names) > 0:
    raise ValueError(
        f"In pmap with axis names {axis_names}, but wrong axis name "
        f"({axis_name}) was provided. This is likely a bug."
    )

  return False


def wrap_if_pmap(
    p_func: Callable[[TArrayTree, str], TArrayTree],
) -> Callable[[TArrayTree, str | None], TArrayTree]:
  """Wraps `p_func` to be executed only when inside a `jax.pmap` context."""

  @functools.wraps(p_func)
  def p_func_if_pmap(obj: TArrayTree, axis_name: str | None) -> TArrayTree:
    return p_func(obj, axis_name) if in_pmap(axis_name) else obj

  return p_func_if_pmap


# TODO(jamesmartens,botev): We no longer use wrap_if_pmap in the below
# definitions since it doesn't seem to transmit type info properly. Investigate?
def pmean_if_pmap(obj: TArrayTree, axis_name: str | None) -> TArrayTree:
  return lax.pmean(obj, axis_name) if in_pmap(axis_name) else obj


def psum_if_pmap(obj: TArrayTree, axis_name: str | None) -> TArrayTree:
  return lax.psum(obj, axis_name) if in_pmap(axis_name) else obj


pmap_mean = jax.pmap(lambda x: lax.pmean(x, "i"), axis_name="i")
pmap_sum = jax.pmap(lambda x: lax.psum(x, "i"), axis_name="i")


def is_scalar(x: Any) -> bool:
  return isinstance(x, (float, int)) or (
      isinstance(x, jax.Array) and not x.shape
  )


def using_legacy_pmap() -> bool:
  """Returns whether the legacy pmap is being used."""
  return False


def get_device_n_contents(obj: TArrayTree, n: int) -> TArrayTree:
  """Gets the contents from pmap output for device n."""

  def _get_device_n_contents(value: Numeric) -> Numeric:

    if is_scalar(value):
      return value

    if using_legacy_pmap():
      return value[n]

    assert isinstance(value, jax.Array)

    if isinstance(value.sharding, jax.sharding.SingleDeviceSharding):
      return value[n]

    assert isinstance(value.sharding, jax.NamedSharding)

    shard_data = value.addressable_shards[n].data
    if value.sharding.spec[0] is None:
      return shard_data

    return shard_data.squeeze(0)

  return jax.tree_util.tree_map(_get_device_n_contents, obj)


def get_first(obj: TArrayTree) -> TArrayTree:
  return get_device_n_contents(obj, 0)


def get_mean(obj: TArrayTree) -> TArrayTree:
  """Returns the average of `obj` over different devices."""
  return get_first(pmap_mean(obj))


def get_sum(obj: TArrayTree) -> TArrayTree:
  """Returns the sum of `obj` over different devices."""
  return get_first(pmap_sum(obj))


_broadcast_all_local_devices_legacy = jax.pmap(lambda x: x)
_broadcast_all_local_devices_cache: dict[
    str | None, Callable[[TArrayTree], TArrayTree]
] = {}


def broadcast_all_local_devices(
    obj: TArrayTree, axis_name: str | None = None
) -> TArrayTree:
  """Broadcasts `obj` to all local Jax devices.

  Args:
    obj: A pytree to broadcast.
    axis_name: Optional axis name for the pmap.

  Returns:
    The broadcasted pytree.
  """
  if types.tree_is_empty(obj):
    return obj

  # When no axis_name provided, use legacy pmap.
  if axis_name is None:
    return _broadcast_all_local_devices_legacy(obj)

  devices = jax.local_devices()
  mesh = jax.sharding.Mesh(devices, (axis_name,))
  sharding = jax.NamedSharding(mesh, jax.sharding.PartitionSpec(axis_name))

  def _broadcast_with_axis(x):
    return jax.device_put(x, sharding)

  return jax.tree_util.tree_map(_broadcast_with_axis, obj)


pmap_zeros_like = jax.pmap(lambda x: jax.tree_util.tree_map(jnp.zeros_like, x))
jit_zeros_like = jax.jit(lambda x: jax.tree_util.tree_map(jnp.zeros_like, x))


def replicate_all_local_devices(
    obj: TArrayTree, axis_name: str | None = None
) -> TArrayTree:
  """Replicates `obj` to all local Jax devices.

  Args:
    obj: A pytree to replicate.
    axis_name: Optional axis name for sharding. When the result will be passed
      to a pmap with a specific axis_name, this should match to avoid mesh
      sharding mismatches.

  Returns:
    The replicated pytree.
  """
  if types.tree_is_empty(obj):
    return obj

  devices = jax.local_devices()

  # When no axis_name is provided, use the original device_put_replicated.
  if axis_name is None:
    return jax.device_put_replicated(obj, devices=devices)

  mesh = jax.sharding.Mesh(devices, (axis_name,))
  sharding = jax.NamedSharding(mesh, jax.P(axis_name))

  def _replicate_with_axis(x):
    # Stack to add the device dimension, then device_put with sharding.
    stacked = jnp.stack([x] * len(devices))
    return jax.device_put(stacked, sharding)

  return jax.tree_util.tree_map(_replicate_with_axis, obj)


def make_different_rng_key_on_all_devices(rng: PRNGKey) -> PRNGKey:
  """Makes a different PRNG for all Jax devices and processes."""

  rng = jax.random.fold_in(rng, jax.process_index())
  rng = jax.random.split(rng, jax.local_device_count())

  return broadcast_all_local_devices(rng)


p_split = jax.pmap(lambda key: tuple(jax.random.split(key)))

p_split_num = jax.pmap(lambda key, num: tuple(jax.random.split(key, num)),
                       static_broadcasted_argnums=1)


default_device_sync = None


def host_sync(
    obj: TArrayTree,
    sync_op: Callable[[TArrayTree, str], TArrayTree],
) -> TArrayTree:
  """Syncs `obj` across multiple hosts with the operation `sync_op`."""

  # The implementation here is to use the pmap syncing mechanisms but with only
  # the default device of each host. Technically we could do this with all
  # the devices on each host, but that would possibly be wasteful.

  if jax.process_count() > 1:

    # We set default_device_sync here because calling jax.local_devices during
    # the library import stage will break JAX.

    global default_device_sync

    if default_device_sync is None:

      default_devices = [jax.local_devices(process_index=p_idx)[0]
                         for p_idx in range(jax.process_count())]

      default_device_sync = jax.pmap(lambda x, sync_op: sync_op(x, "i"),
                                     devices=default_devices,
                                     axis_name="i",
                                     static_broadcasted_argnums=1)

    obj = jax.tree_util.tree_map(lambda x: jnp.expand_dims(x, axis=0), obj)

    return get_first(default_device_sync(obj, sync_op))

  return obj


def host_all_gather(x: TArrayTree) -> TArrayTree:
  """Gathers on every host the values of the PyTree leaves `x`."""
  return host_sync(x, lax.all_gather)


def host_mean(x: TArrayTree) -> TArrayTree:
  """Computes the mean of the PyTree leaves of `x` over multiple hosts."""
  return host_sync(x, lax.pmean)


def sync_and_divide_value(
    value: TArrayTree,
    counter: Numeric,
    axis_name: str | None = None,
) -> TArrayTree:
  """Computes the mean of `value` over all hosts and divides it by `counter`."""
  value = jax.tree_util.tree_map(lambda x: x / counter, value)
  return pmean_if_pmap(value, axis_name)


jit_sync_and_divide_value = jax.jit(sync_and_divide_value)
pmap_sync_and_divide_value = jax.pmap(
    functools.partial(sync_and_divide_value, axis_name="i"),
    axis_name="i",
)


# We might be able to change this to "return jnp.array(x)" in newer JAX
# versions. Or maybe we can use jnp.copy now?
def copy_array(x: Array) -> Array:
  """Copies a Jax array so that it can be donated freely."""
  return x + jnp.zeros_like(x)


copy_obj = jax.jit(lambda x: jax.tree_util.tree_map(copy_array, x))
_pmap_copy_obj = jax.pmap(copy_obj)


def pmap_copy_obj(x: TArrayTree | None) -> TArrayTree | None:

  # pmap will fail to work if passed a totally empty tree
  if x is None:
    return None

  if types.tree_is_empty(x):
    # this does a shallow copy of the tree similar to .copy():
    (flattened, structure) = jax.tree_util.tree_flatten(x)
    return jax.tree_util.tree_unflatten(structure, flattened)

  return _pmap_copy_obj(x)


def distribute_thunks(
    thunks: Sequence[Callable[[], TArrayTree]],
    pmap_axis_name: str,
    ) -> TArrayTree:
  """Distributes the computation of a list of thunks over the pmapped devices.

  Given a list of thunks, this function distributes their computation over the
  devices of the current pmap in a round-robin fashion, syncronizes the results
  across devices, and then returns them as a sequence of PyTrees.

  Note that this function is meant to be used in a compiled context, and may
  call ``thunk[i]()`` several times for each i, with all but one call getting
  "optimized away" by XLA.

  Args:
    thunks: A sequence of callables performing the desired computations. Each
      callable must take zero arguments and return a PyTree of JAX arrays. As
      with callables passed to (most) standard JAX API functions, these need to
      be stateless and free of side effects. The output of each callable must be
      the same regardless of the device it is executed on.
    pmap_axis_name: The name of the pmap axis to use.

  Returns:
    A sequence of PyTrees that are the output of the corresponding element of
    ``thunks``.
  """

  # The strategy here is to make a callable for each device which executes only
  # the thunks i such that i % total_devices == device_index, and returns a tree
  # of zeros for the remaining thunks. We then do a lax.switch over these based
  # on device_index, and return psum over these. Note that the more obvious way
  # of doing this, which is to perform a psum over the output of a sequence of
  # lax.cond calls (with one for each thunk), won't work in general. This is
  # because in order to save memory, XLA will sometimes elect to execute these
  # conds sequentially instead of in parallel.

  if not in_pmap(pmap_axis_name):
    raise ValueError(f"Provided pmap_axis_name {pmap_axis_name} is not a valid "
                     "pmap axis in current pmap (or this function was not "
                     "called in a pmap).")

  assert pmap_axis_name is not None

  axis_names = _axis_name_tuple(pmap_axis_name)
  total_devices = lax.psum(1, axis_name=pmap_axis_name)  # returns a constant
  if len(axis_names) == 1:
    current_device_index = lax.axis_index(axis_names[0])
  else:
    # Linearise the multi-axis shard_map index so distributed thunk work is
    # spread over the full data mesh, not just one named axis.
    current_device_index = 0
    stride = 1
    for axis in reversed(axis_names):
      current_device_index = current_device_index + lax.axis_index(axis) * stride
      stride = stride * lax.psum(1, axis_name=axis)

  # This should get optimized away by XLA since we don't use the values:
  dummy_output_trees = tuple(thunk() for thunk in thunks)

  def make_branch(device_index):

    def branch():
      """Execute only thunks i such that i % total_devices == device_index."""

      outs = []
      for i in range(len(thunks)):

        if i % total_devices == device_index:
          outs.append(thunks[i]())
        else:
          outs.append(
              jax.tree_util.tree_map(jnp.zeros_like, dummy_output_trees[i]))

      return tuple(outs)

    return branch

  branches = tuple(make_branch(device_index)
                   for device_index in range(total_devices))

  output_trees = jax.lax.switch(current_device_index, branches)

  return jax.lax.psum(output_trees, axis_name=pmap_axis_name)