PR #26906: [jax.distributed] Allow explicitly setting slice_index
Imported from GitHub PR https://github.com/jax-ml/jax/pull/26906
Allows overriding the slice index used by XLA.
More explicit control over which slice a device ends up in is desirable:
- Various parts of the ecosystem equate slices with "devices communicating via fast interconnect". With the arrival of NVL72 we want devices managed by multiple hosts to form a single slice.
- For debugging purposes it can be useful to allow devices on the same host (managed in separate processes) to be treated as different slices. For example, [Orbax](https://github.com/google/orbax)'s local checkpointing presumes the existence of at least two slices, so overriding the boot id will allow us to test local checkpointing on a single host.
(Companion PR in XLA: https://github.com/openxla/xla/pull/23347)
Copybara import of the project:
--
45aa7ce316bb05ebcc3f3ed2d888385923285e58 by Georg Stefan Schmid <gschmid@nvidia.com>:
[jax.distributed] Allow overriding XLA slice_index
Merging this change closes #26906
COPYBARA_INTEGRATE_REVIEW=https://github.com/jax-ml/jax/pull/26906 from gspschmid:gschmid/jax-override-boot-id 45aa7ce316bb05ebcc3f3ed2d888385923285e58
PiperOrigin-RevId: 744012253