Skip to content

Add single_host mode to run on one host of a multi-host slice - #141

Open
simrankaurb wants to merge 1 commit into
AI-Hypercomputer:chsfrom
simrankaurb:chs-single-host-mode
Open

simrankaurb wants to merge 1 commit into
AI-Hypercomputer:chsfrom
simrankaurb:chs-single-host-mode

Conversation

@simrankaurb

Copy link
Copy Markdown

What

Adds a top-level single_host config key to Ironwood/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.

single_host: true
benchmarks:
- benchmark_name: gemm_multiple_run
  benchmark_sweep_params:
  - {m: 16384, k: 16384, n: 16384, num_runs: 100, dtype: 'bfloat16', run_on_local_node: True}

Why

Cluster Health Scanner runs gemm_multiple_run, single_device_hbm_copy, host_device and 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 to LIBTPU_INIT_ARGS.
    • set_libtpu_init_args() sets LIBTPU_INIT_ARGS and keeps the single-host flags when the mode is on.
    • Why this works: libtpu applies LIBTPU_INIT_ARGS after the flags it derives from the environment, so these win. With host bounds 1,1,1 libtpu uses task id 0 and a localhost-only slice builder, and ignores the other TPU_WORKER_HOSTNAMES. Chips-per-host bounds already describe one host, so they are left unchanged.
  • run_benchmark.py:
    • Reads single_host, which must be a bool and can't be combined with --multithreaded.
    • Calls enable_single_host_mode() right after the benchmark module is imported, because modules overwrite LIBTPU_INIT_ARGS at import.
    • After each benchmark run, checks that the backend really is single-host (device_count == local_device_count).
  • benchmark_collectives.py: the four functions that set LIBTPU_INIT_ARGS at run time now call set_libtpu_init_args(), so single-node collectives keep the single-host flags.
  • README: documents the key.

When single_host is unset, behaviour is unchanged.

Relation to other PRs

Testing

  • py_compile passes.
  • Offline unit test with stubbed jax/ray/pandas, 7/7 passing. It checks that:
    • the flags survive module-import overwrites and function-level overwrites (collectives style);
    • the flags come last and appear once;
    • behaviour is unchanged when single_host is off;
    • a multi-host backend is rejected;
    • a non-bool value and --multithreaded are rejected.
  • Not yet validated on TPU hardware. Next step: run it end to end on a multi-host tpu7x / v6e slice through the companion CHS change.

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.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread Ironwood/src/common.py
Comment on lines +34 to +41
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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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.

Suggested change
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)

Comment on lines 428 to +432
except Exception as e: # pylint: disable=broad-except
print(f"Benchmark func failed: {e}")
continue
if single_host:
check_single_host_backend()

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

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()

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant