Download src/kfac_jax/_src/utils/parallel.py from Twobombs/HamiltonZero: direct link, hf CLI and curl.
- Browser
- Download file 13 kB
-
https://huggingface.co/Twobombs/HamiltonZero/resolve/main/src/kfac_jax/_src/utils/parallel.py
- Command line
-
hf download hf://Twobombs/HamiltonZero/src/kfac_jax/_src/utils/parallel.py
-
curl -L -o parallel.py https://huggingface.co/Twobombs/HamiltonZero/resolve/main/src/kfac_jax/_src/utils/parallel.py
13 kB
| # 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.""" | |
| 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) | |