Repository navigation
Add single_host mode to run on one host of a multi-host slice - #141
simrankaurb wants to merge 1 commit into
Conversation
On a multi-host TPU slice, libtpu gets the slice-wide topology from GKE / the TPU VM metadata and waits for every host to join before the TPU backend comes up, so a benchmark cannot run on a single host by itself. Node-level health checks (e.g. Cluster Health Scanner's tpu-gemm, tpu-hbm, tpu-host-device and single-node collectives tests) need to benchmark one host at a time. Add a top-level 'single_host' config key. When set, run_benchmark appends --deepsea_host_bounds=1,1,1 (plus no wrap/twist) to LIBTPU_INIT_ARGS after the benchmark module is imported. libtpu applies LIBTPU_INIT_ARGS after its environment-derived flags, so the host comes up as a standalone 1-host slice: task id 0, localhost-only slice builder, and the other hosts in TPU_WORKER_HOSTNAMES are ignored. Benchmarks that set LIBTPU_INIT_ARGS inside the benchmark function (benchmark_collectives) now go through common.set_libtpu_init_args(), which keeps the single-host flags when the mode is on. The run fails if the backend still reports devices on other hosts.
There was a problem hiding this comment.
Code Review
This pull request introduces support for running benchmarks on a single host of a multi-host slice by adding a single_host configuration option, updating the documentation, and providing helper functions in common.py to manage single-host LIBTPU_INIT_ARGS flags. Feedback on the changes suggests making set_libtpu_init_args more robust by handling string inputs and filtering out conflicting flags, as well as checking the backend configuration inside the exception handler in run_benchmark.py to ensure the script fails early and clearly if single-host initialization fails.
| libtpu_init_args = list(libtpu_init_args) | ||
| if _single_host_mode_enabled: | ||
| libtpu_init_args += [ | ||
| arg | ||
| for arg in SINGLE_HOST_LIBTPU_ARGS | ||
| if arg not in libtpu_init_args | ||
| ] | ||
| os.environ["LIBTPU_INIT_ARGS"] = " ".join(libtpu_init_args) |
There was a problem hiding this comment.
To make set_libtpu_init_args more robust and defensive, we should handle cases where libtpu_init_args is passed as a single string instead of an iterable of strings (which would otherwise be split into individual characters). Additionally, when _single_host_mode_enabled is active, we should filter out any existing conflicting flags (e.g., --deepsea_host_bounds, --deepsea_wrap, --deepsea_twist) from the input arguments before appending the single-host flags to avoid duplicate or conflicting arguments in LIBTPU_INIT_ARGS.
| libtpu_init_args = list(libtpu_init_args) | |
| if _single_host_mode_enabled: | |
| libtpu_init_args += [ | |
| arg | |
| for arg in SINGLE_HOST_LIBTPU_ARGS | |
| if arg not in libtpu_init_args | |
| ] | |
| os.environ["LIBTPU_INIT_ARGS"] = " ".join(libtpu_init_args) | |
| if isinstance(libtpu_init_args, str): | |
| libtpu_init_args = libtpu_init_args.split() | |
| else: | |
| libtpu_init_args = list(libtpu_init_args) | |
| if _single_host_mode_enabled: | |
| prefixes = tuple(arg.split("=")[0] + "=" for arg in SINGLE_HOST_LIBTPU_ARGS) | |
| libtpu_init_args = [arg for arg in libtpu_init_args if not arg.startswith(prefixes)] | |
| libtpu_init_args.extend(SINGLE_HOST_LIBTPU_ARGS) | |
| os.environ["LIBTPU_INIT_ARGS"] = " ".join(libtpu_init_args) |
| except Exception as e: # pylint: disable=broad-except | ||
| print(f"Benchmark func failed: {e}") | ||
| continue | ||
| if single_host: | ||
| check_single_host_backend() |
There was a problem hiding this comment.
If benchmark_func fails with an exception (for example, due to a misconfigured multi-host backend or a timeout during initialization), the current code will print the error and silently continue to the next parameter sweep. If single_host mode is enabled, we should check the backend configuration immediately in the except block to raise a clear RuntimeError and terminate the script, rather than continuing to subsequent iterations which will also fail or hang.
except Exception as e: # pylint: disable=broad-except
if single_host:
check_single_host_backend()
print(f"Benchmark func failed: {e}")
continue
if single_host:
check_single_host_backend()
What
Adds a top-level
single_hostconfig key toIronwood/src/run_benchmark.py. When it is set, the benchmark runs on one host of a multi-host TPU slice by itself: that host comes up as a standalone 1-host slice using only its local chips.Why
Cluster Health Scanner runs
gemm_multiple_run,single_device_hbm_copy,host_deviceand single-node collectives as per-node health checks. On a multi-host slice (e.g. tpu7x 4x4x4, v6e 4x4), libtpu reads the slice-wide topology from GKE / TPU VM metadata. It then waits in slice formation until every host joins, so today the whole slice must be idle and all of its hosts must run the benchmark together.run_on_local_node/jax.distributed(#138) don't change this: the wait happens inside libtpu when the backend initializes.How
common.py:SINGLE_HOST_LIBTPU_ARGS = --deepsea_host_bounds=1,1,1 --deepsea_wrap=false,false,false --deepsea_twist=false.enable_single_host_mode()appends these flags toLIBTPU_INIT_ARGS.set_libtpu_init_args()setsLIBTPU_INIT_ARGSand keeps the single-host flags when the mode is on.LIBTPU_INIT_ARGSafter the flags it derives from the environment, so these win. With host bounds1,1,1libtpu uses task id 0 and a localhost-only slice builder, and ignores the otherTPU_WORKER_HOSTNAMES. Chips-per-host bounds already describe one host, so they are left unchanged.run_benchmark.py:single_host, which must be a bool and can't be combined with--multithreaded.enable_single_host_mode()right after the benchmark module is imported, because modules overwriteLIBTPU_INIT_ARGSat import.device_count == local_device_count).benchmark_collectives.py: the four functions that setLIBTPU_INIT_ARGSat run time now callset_libtpu_init_args(), so single-node collectives keep the single-host flags.When
single_hostis unset, behaviour is unchanged.Relation to other PRs
jax.distributed.initialize) is complementary: that PR covers the JAX coordination layer, and this one covers libtpu slice formation. If both land,single_host: trueshould also skipjax.distributed.initialize().gemm_multiple_run) is independent.Testing
py_compilepasses.single_hostis off;--multithreadedare rejected.