diff --git a/.gitignore b/.gitignore index c824efdf..9dde930c 100644 --- a/.gitignore +++ b/.gitignore @@ -101,6 +101,9 @@ instance/ # Sphinx documentation docs/_build/ +# Local agent planning artifacts +docs/superpowers/ + # Jupyter Notebook .ipynb_checkpoints diff --git a/CHANGELOG.md b/CHANGELOG.md index 062c6be4..81cb321a 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **Spur (Crusoe) scheduler support** (#157): Spur ships SLURM-compatible CLI shims, but `srun` cannot fan out across nodes (`SLURM_PROCID` is empty), `scontrol show hostname[s]` is unsupported, and the Raft-based control plane makes `squeue`/`sacct` eventually consistent. A new `SpurDeployment` backend reuses the SLURM template and presets but drives multi-node runs with a job **array** of single-node tasks that self-form the cluster through a shared-filesystem rendezvous, and detects completion from per-rank marker files instead of `sacct`. Select it with `"slurm": {"scheduler": "spur"}` — the `slurm` block is otherwise unchanged, so spur is reachable from both `madengine build` and `madengine run --additional-context`. Peers that do not see rank 0's `MASTER_ADDR` within `slurm.rendezvous_timeout` (default 900s) fail with a diagnostic instead of starting with an empty address. The `srun`-based node health preflight is skipped on spur (it cannot target a node there, and the `--nodelist` it pins conflicts with the job array), an explicit `"deploy": "spur"` selects the backend at run time as well as at build time, and monitoring gives up with a diagnostic if the array never appears in `squeue` instead of polling forever. Stock SLURM behaviour is unchanged. + ### Docs - **README rewritten as a concise landing page** (#161): Trimmed the root README from 707 to 258 lines by moving deep reference material (profiling tables, extended config/usage recipes, tips) into `docs/` and linking out. Replaced the stale ASCII architecture block and unreferenced `docs/img` PNGs with accurate inline Mermaid diagrams for the layered architecture, build→run→report pipeline, and deployment-target inference; added matching diagrams to `docs/deployment.md` and `docs/README.md`. Also corrects numerous stale references across docs: `--csv-file` → `--csv-file-path`/`--file`, missing `database` command flags (`--unique-key`/`-k`, `--batch-size`, `--no-upsert`, `--no-index`, `--dry-run`, `MONGO_AUTH_SOURCE`/`MONGO_TIMEOUT_MS`), wrong `run --output`/`--tools-config` defaults, `megatron` → `megatron-lm` launcher name, fabricated `timeout_multiplier`/`service_account` config keys, missing Kubernetes/SLURM `additional_context` keys, `DOCKER_CONFIG`/`MAD_SKIP_DOCKER_LOGIN` documentation, and corrected SGLang Disaggregated minimum node counts/split formula for SLURM vs. Kubernetes. diff --git a/docs/deployment.md b/docs/deployment.md index c913b117..95d0a2b0 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -249,9 +249,46 @@ The deployment target is automatically detected from the `slurm` key in the conf - `reservation`: SLURM reservation name; forwarded to srun health/cleanup commands - `time`: Wall time limit (HH:MM:SS) - `exclusive`: Exclusive node access (default: `true`) +- `scheduler`: SLURM flavor - `slurm` (default) or `spur` (see below) +- `rendezvous_timeout`: spur only; seconds a node waits for rank 0 to publish `MASTER_ADDR` (default: 900) See [examples/slurm-configs/](../examples/slurm-configs/) for complete examples. +### Spur (Crusoe) Scheduler + +Spur exposes SLURM-compatible CLI shims but `srun` cannot fan tasks out across +nodes, so multi-node runs use a job **array** of single-node tasks instead: each +array task runs on one node, `SLURM_ARRAY_TASK_ID` is the node rank, and the +tasks self-form the cluster through a shared-filesystem rendezvous (rank 0 +publishes its transport IP; the other ranks read it as `MASTER_ADDR`). + +Select it with `slurm.scheduler`; everything else in the `slurm` block is +unchanged: + +```json +{ + "slurm": { + "scheduler": "spur", + "partition": "gpu", + "nodes": 4, + "gpus_per_node": 8, + "time": "02:00:00" + } +} +``` + +A job array carries no gang-scheduling guarantee, so tasks may start minutes +apart. A node that does not see rank 0's address within `rendezvous_timeout` +fails the run with a diagnostic rather than starting with an empty +`MASTER_ADDR`; raise the timeout if your queue wait is longer than that. + +`slurm.output_dir` must be on a filesystem shared by every node - it holds the +rendezvous files. + +The `srun`-based node health preflight (`enable_node_check`) is skipped on spur: +`srun -w ` does not run on the requested node there, and the `--nodelist` +it would pin conflicts with the job array, whose tasks each request one node. + ### Multi-Node Training For distributed training across SLURM nodes: diff --git a/docs/superpowers/plans/2026-08-27-pinned-image-digest.md b/docs/superpowers/plans/2026-08-27-pinned-image-digest.md deleted file mode 100644 index dad4b9d3..00000000 --- a/docs/superpowers/plans/2026-08-27-pinned-image-digest.md +++ /dev/null @@ -1,1623 +0,0 @@ -# Pinned Image Digest Implementation Plan - -> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. - -**Goal:** Record the real digest of every image madengine pushes, and let users opt into pinning run-time pulls to that digest so a moved registry tag fails loudly instead of silently running the wrong image. - -**Architecture:** A new leaf module `madengine/core/image_digest.py` holds four small pure functions (two parsers, one reference builder, one policy resolver that raises when enforcement is on but a digest is missing). `DockerBuilder.push_image()` records digests into a `self.pushed_digests` dict — mirroring the existing `self.built_images` / `self.built_models` pattern — and the two push call sites copy the digest into `build_info["image_digest"]`, which rides into `build_manifest.json` unchanged. At run time a `--require-pinned-image` flag (mirrored as the `require_pinned_image` additional-context key) flows through `RunOrchestrator.additional_context` into all three execution paths, each calling the same `resolve_pinned_image()` helper. - -**Tech Stack:** Python 3, Typer (CLI), pytest + `unittest.mock`, Jinja2 (SLURM/K8s templates), Docker CLI. - ---- - -## Background you need before starting - -Read the design spec first: `docs/superpowers/specs/2026-08-27-pinned-image-digest-design.md`. - -Facts about this codebase that the tasks below depend on: - -- `Console.sh(command)` (`src/madengine/core/console.py:138`) runs a shell command and **returns its stdout as a stripped string**, raising `RuntimeError` on non-zero exit unless `canFail=True`. This is how we capture `docker push` output. -- `build_info` dicts are created in `DockerBuilder.build_image()` (`src/madengine/execution/docker_builder.py:311-322`) and serialized wholesale into `build_manifest.json` under `built_images` by `export_build_manifest()`. **Any new key added to `build_info` appears in the manifest automatically** — no serializer changes needed. -- There are exactly two places that call `push_image` and set `build_info["registry_image"]`: `docker_builder.py:769` (single-arch) and `docker_builder.py:977` (per-GPU-arch). Both must record the digest. -- `ContainerRunner.run_models_from_manifest()` merges the manifest's `context` dict over its own `additional_context` (`src/madengine/execution/container_runner.py:2799-2800`). `RunOrchestrator._load_and_merge_manifest()` writes selected runtime keys into `manifest["context"]` (`src/madengine/orchestration/run_orchestrator.py:499-503`). Adding our key to that merge list is what makes the flag survive into the nested `madengine run` that the standard SLURM job script executes on each compute node. -- `_docker_image_ref_for_log_naming()` (`src/madengine/execution/container_runner_helpers.py:232`) already strips `@sha256:...`, so pinned references produce the same log filenames as tags. Task 9 locks that in with a test. -- Test style in this repo: `pytest` classes named `TestX` with plain `assert`, `MagicMock` for `Context`/`Console`, `tmp_path` for manifests. Follow the surrounding file's style in each test file you touch. - -### Deliberate deviations from the spec (do not "fix" these) - -1. **SLURM enforcement point.** The spec's table says the pinned reference goes into the generated `srun docker pull` line in `slurm.py:544`. Implementing only that would pull the pinned image and then `docker run` the *tag* — the exact race the feature exists to close. Instead we set `env_vars["DOCKER_IMAGE_NAME"]` to the pinned reference (Task 8); the pull line interpolates that same variable, so both pull and run are pinned by one change. -2. **Digest capture is not added to the build-on-compute-node path** (`build_orchestrator._execute_build_on_compute` at `build_orchestrator.py:748`, which pushes from a generated bash script inside an sbatch job and writes `built_images` entries at `:1224`/`:1239` with no `registry_image` and no digest). Manifests from that path have no `image_digest`, so `--require-pinned-image` runs against them fail fast with the Task 2 error. This is documented in Task 10, not implemented. - ---- - -## File Structure - -**Create:** - -| File | Responsibility | -|---|---| -| `src/madengine/core/image_digest.py` | Pure helpers: parse a digest out of `docker push` / `docker image inspect` output, build a `repo@sha256:...` reference, and apply the enforcement policy (return-as-is / pin / raise). No I/O, no Docker calls. | -| `tests/unit/test_image_digest.py` | Unit tests for all four functions in the module above. | - -**Modify:** - -| File | Change | -|---|---| -| `src/madengine/execution/docker_builder.py` | Capture the pushed digest in `push_image()`; copy it into `build_info["image_digest"]` at both push call sites. | -| `src/madengine/cli/commands/run.py` | Add the `--require-pinned-image` flag; pass it into both `create_args_namespace(...)` calls. | -| `src/madengine/orchestration/run_orchestrator.py` | Fold the flag into `additional_context`; persist it into `manifest["context"]`. | -| `src/madengine/execution/container_runner.py` | Resolve the pinned reference before the registry pull. | -| `src/madengine/deployment/k8s_template_context.py` | Resolve the pinned reference for the pod spec `image` field. | -| `src/madengine/deployment/slurm.py` | Resolve the pinned reference for `DOCKER_IMAGE_NAME` in the slurm_multi wrapper. | -| `docs/cli-reference.md`, `docs/configuration.md` | Document the flag and the context key. | - -**Test files touched:** `tests/unit/test_image_digest.py` (new), `tests/unit/test_docker_builder.py`, `tests/unit/test_orchestration.py`, `tests/unit/test_container_runner.py`, `tests/unit/test_k8s.py`, `tests/unit/test_slurm_multi.py`, `tests/unit/test_execution.py`. - ---- - -## Task 1: Digest parsing and pinned-reference construction - -**Files:** -- Create: `src/madengine/core/image_digest.py` -- Test: `tests/unit/test_image_digest.py` - -- [ ] **Step 1: Write the failing tests** - -Create `tests/unit/test_image_digest.py`: - -```python -"""Unit tests for madengine.core.image_digest. - -Covers digest extraction from `docker push` / `docker image inspect` output and -construction of pinned `repo@sha256:...` references. - -Copyright (c) Advanced Micro Devices, Inc. All rights reserved. -""" - -import pytest - -from madengine.core.image_digest import ( - build_pinned_reference, - parse_push_digest, - parse_repo_digest, -) - - -DIGEST = "sha256:" + "df36ef7e" * 8 # 64 hex chars -OTHER_DIGEST = "sha256:" + "cbb7e5ed" * 8 - - -class TestParsePushDigest: - """parse_push_digest extracts the digest line emitted by `docker push`.""" - - def test_typical_push_output(self): - output = ( - "The push refers to repository [docker.io/myorg/ci]\n" - "a1b2c3d4e5f6: Pushed\n" - "9f8e7d6c5b4a: Layer already exists\n" - f"mymodel: digest: {DIGEST} size: 4738\n" - ) - assert parse_push_digest(output) == DIGEST - - def test_no_space_after_colon(self): - assert parse_push_digest(f"mymodel: digest:{DIGEST} size: 12") == DIGEST - - def test_last_digest_wins_when_multiple(self): - output = ( - f"tag-a: digest: {OTHER_DIGEST} size: 10\n" - f"tag-b: digest: {DIGEST} size: 10\n" - ) - assert parse_push_digest(output) == DIGEST - - def test_uppercase_hex_is_not_matched(self): - assert parse_push_digest("mymodel: digest: sha256:ABC size: 1") is None - - def test_output_without_digest_line(self): - assert parse_push_digest("The push refers to repository [x]\nlayer: Pushed\n") is None - - def test_empty_output(self): - assert parse_push_digest("") is None - - def test_none_output(self): - assert parse_push_digest(None) is None - - -class TestParseRepoDigest: - """parse_repo_digest extracts the digest from a `repo@sha256:...` reference.""" - - def test_repo_digest_reference(self): - assert parse_repo_digest(f"myorg/ci@{DIGEST}") == DIGEST - - def test_registry_with_port(self): - assert parse_repo_digest(f"localhost:5000/myorg/ci@{DIGEST}") == DIGEST - - def test_surrounding_whitespace_and_quotes(self): - assert parse_repo_digest(f" 'myorg/ci@{DIGEST}' \n") == DIGEST - - def test_empty_repodigests_placeholder(self): - # `docker image inspect` prints this when RepoDigests is empty. - assert parse_repo_digest("") is None - - def test_none_output(self): - assert parse_repo_digest(None) is None - - -class TestBuildPinnedReference: - """build_pinned_reference produces repo@sha256:... with any tag stripped.""" - - def test_strips_tag(self): - assert build_pinned_reference(f"myorg/ci:mymodel", DIGEST) == f"myorg/ci@{DIGEST}" - - def test_no_tag(self): - assert build_pinned_reference("myorg/ci", DIGEST) == f"myorg/ci@{DIGEST}" - - def test_registry_port_is_not_mistaken_for_a_tag(self): - out = build_pinned_reference("localhost:5000/myorg/ci:latest", DIGEST) - assert out == f"localhost:5000/myorg/ci@{DIGEST}" - - def test_registry_port_without_tag(self): - out = build_pinned_reference("localhost:5000/myorg/ci", DIGEST) - assert out == f"localhost:5000/myorg/ci@{DIGEST}" - - def test_bare_name_with_tag(self): - assert build_pinned_reference("ci-dummy:latest", DIGEST) == f"ci-dummy@{DIGEST}" - - def test_existing_digest_is_replaced(self): - out = build_pinned_reference(f"myorg/ci@{OTHER_DIGEST}", DIGEST) - assert out == f"myorg/ci@{DIGEST}" -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_image_digest.py -v` -Expected: FAIL — `ModuleNotFoundError: No module named 'madengine.core.image_digest'` - -- [ ] **Step 3: Write the implementation** - -Create `src/madengine/core/image_digest.py`: - -```python -#!/usr/bin/env python3 -""" -Registry image digest helpers for madengine. - -Build pushes record the digest of the image they push; runs can optionally be -pinned to that digest so a registry tag that moved between build and run fails -loudly instead of silently resolving to a different image. - -Copyright (c) Advanced Micro Devices, Inc. All rights reserved. -""" - -import re -import typing - - -# `docker push` prints e.g. "mytag: digest: sha256:<64 hex> size: 4738". -_PUSH_DIGEST_RE = re.compile(r"digest:\s*(sha256:[0-9a-f]{64})") - -# `docker image inspect --format '{{index .RepoDigests 0}}'` prints "repo@sha256:<64 hex>". -_REPO_DIGEST_RE = re.compile(r"@(sha256:[0-9a-f]{64})") - - -def parse_push_digest(push_output: typing.Optional[str]) -> typing.Optional[str]: - """Extract the pushed image digest from `docker push` output. - - Args: - push_output: Combined stdout/stderr of the push command. - - Returns: - The digest (``sha256:...``), or None if no digest line was present. - When several digest lines appear the last one is returned, which is the - digest of the reference the push command was invoked with. - """ - if not push_output: - return None - matches = _PUSH_DIGEST_RE.findall(push_output) - return matches[-1] if matches else None - - -def parse_repo_digest(inspect_output: typing.Optional[str]) -> typing.Optional[str]: - """Extract the digest from a ``repo@sha256:...`` reference. - - Args: - inspect_output: Output of - ``docker image inspect --format '{{index .RepoDigests 0}}' ``. - - Returns: - The digest (``sha256:...``), or None if the output held no digest - (e.g. ```` when the image has no RepoDigests entry). - """ - if not inspect_output: - return None - match = _REPO_DIGEST_RE.search(inspect_output) - return match.group(1) if match else None - - -def build_pinned_reference(registry_image: str, digest: str) -> str: - """Build a digest-pinned image reference. - - Any existing tag or digest on ``registry_image`` is dropped. A port in the - registry host (``localhost:5000/org/img``) is not mistaken for a tag because - only the final path segment is inspected for ``:``. - - Args: - registry_image: Image reference, with or without tag/digest. - digest: Digest to pin to (``sha256:...``). - - Returns: - A reference of the form ``repo@sha256:...``. - """ - repo = registry_image.split("@", 1)[0] - last_slash = repo.rfind("/") - tail = repo[last_slash + 1 :] - if ":" in tail: - repo = repo[: last_slash + 1] + tail.split(":", 1)[0] - return f"{repo}@{digest}" -``` - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_image_digest.py -v` -Expected: PASS (18 tests) - -- [ ] **Step 5: Commit** - -```bash -git add src/madengine/core/image_digest.py tests/unit/test_image_digest.py -git commit -m "feat(image-digest): add digest parsing and pinned reference helpers" -``` - ---- - -## Task 2: Enforcement policy resolver - -**Files:** -- Modify: `src/madengine/core/image_digest.py` -- Test: `tests/unit/test_image_digest.py` - -`resolve_pinned_image()` is the single place the "pin, pass through, or fail" decision is made. All three execution paths (local Docker, K8s, SLURM) call it so they cannot drift. - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_image_digest.py`: - -```python -class TestResolvePinnedImage: - """resolve_pinned_image applies the --require-pinned-image policy.""" - - def test_disabled_returns_image_unchanged_without_digest(self): - assert resolve_pinned_image("myorg/ci:mymodel", None, False) == "myorg/ci:mymodel" - - def test_disabled_returns_image_unchanged_even_with_digest(self): - # Default behaviour must not change for existing users: the digest rides - # along in the manifest but is never used unless enforcement is on. - assert resolve_pinned_image("myorg/ci:mymodel", DIGEST, False) == "myorg/ci:mymodel" - - def test_enabled_with_digest_returns_pinned_reference(self): - out = resolve_pinned_image("myorg/ci:mymodel", DIGEST, True) - assert out == f"myorg/ci@{DIGEST}" - - def test_enabled_without_digest_raises(self): - with pytest.raises(ConfigurationError): - resolve_pinned_image("myorg/ci:mymodel", None, True, model_name="my_model") - - def test_enabled_with_empty_digest_raises(self): - with pytest.raises(ConfigurationError): - resolve_pinned_image("myorg/ci:mymodel", "", True, model_name="my_model") - - def test_error_message_names_model_and_image(self): - with pytest.raises(ConfigurationError) as excinfo: - resolve_pinned_image("myorg/ci:mymodel", None, True, model_name="my_model") - message = str(excinfo.value) - assert "my_model" in message - assert "myorg/ci:mymodel" in message - - def test_error_carries_actionable_suggestions(self): - with pytest.raises(ConfigurationError) as excinfo: - resolve_pinned_image("myorg/ci:mymodel", None, True, model_name="my_model") - assert excinfo.value.suggestions -``` - -And extend the imports at the top of the file: - -```python -from madengine.core.errors import ConfigurationError -from madengine.core.image_digest import ( - build_pinned_reference, - parse_push_digest, - parse_repo_digest, - resolve_pinned_image, -) -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_image_digest.py -k ResolvePinned -v` -Expected: FAIL — `ImportError: cannot import name 'resolve_pinned_image'` - -- [ ] **Step 3: Write the implementation** - -Add the import near the top of `src/madengine/core/image_digest.py`, below `import typing`: - -```python -from madengine.core.errors import ConfigurationError, create_error_context -``` - -Append to the same file: - -```python -def resolve_pinned_image( - registry_image: str, - image_digest: typing.Optional[str], - require_pinned: bool, - model_name: str = "", -) -> str: - """Resolve the image reference to use for a run. - - Args: - registry_image: Tagged registry reference from the build manifest. - image_digest: Digest recorded at build time, if any. - require_pinned: True when --require-pinned-image / require_pinned_image is set. - model_name: Model name, used only to make the error message actionable. - - Returns: - ``registry_image`` unchanged when enforcement is off, otherwise a - digest-pinned reference. - - Raises: - ConfigurationError: When enforcement is on but the manifest recorded no - digest. Falling back to the tag is deliberately not done: silent - degradation would defeat the guarantee the caller asked for. - """ - if not require_pinned: - return registry_image - - if not image_digest: - raise ConfigurationError( - f"--require-pinned-image is set but the build manifest records no " - f"image digest for model '{model_name}' (image: {registry_image}). " - f"Refusing to pull by tag.", - context=create_error_context( - operation="resolve_pinned_image", - component="image_digest", - additional_info={"model": model_name, "image": registry_image}, - ), - suggestions=[ - "Rebuild with this version of madengine so the push digest is recorded", - "Drop --require-pinned-image / require_pinned_image to pull by tag", - ], - ) - - return build_pinned_reference(registry_image, image_digest) -``` - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_image_digest.py -v` -Expected: PASS (25 tests) - -- [ ] **Step 5: Commit** - -```bash -git add src/madengine/core/image_digest.py tests/unit/test_image_digest.py -git commit -m "feat(image-digest): add resolve_pinned_image enforcement policy" -``` - ---- - -## Task 3: Capture the pushed digest in `DockerBuilder.push_image()` - -**Files:** -- Modify: `src/madengine/execution/docker_builder.py:51`, `:394-403` -- Test: `tests/unit/test_docker_builder.py` - -Digest capture is **always on** and never fails a build. If neither the push output nor `docker image inspect` yields a digest, we log a dim note and move on — the push already succeeded. - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_docker_builder.py`: - -```python -DIGEST = "sha256:" + "df36ef7e" * 8 - - -class TestPushImageRecordsDigest: - """push_image records the pushed image digest in builder.pushed_digests.""" - - def _builder(self, sh_side_effect): - ctx = MagicMock() - ctx.ctx = {} - console = MagicMock() - console.sh = MagicMock(side_effect=sh_side_effect) - builder = DockerBuilder(ctx, console) - builder.rich_console = MagicMock() - return builder - - def test_digest_parsed_from_push_output(self): - def sh(command, *args, **kwargs): - if "docker push" in command: - return f"mymodel: digest: {DIGEST} size: 4738" - return "" - - builder = self._builder(sh) - result = builder.push_image("ci-dummy", "localhost:5000", None, "localhost:5000/ci-dummy") - - assert result == "localhost:5000/ci-dummy" - assert builder.pushed_digests["localhost:5000/ci-dummy"] == DIGEST - - def test_falls_back_to_image_inspect_when_push_output_has_no_digest(self): - def sh(command, *args, **kwargs): - if "docker push" in command: - return "The push refers to repository [localhost:5000/ci-dummy]\nlayer: Pushed" - if "docker image inspect" in command: - return f"localhost:5000/ci-dummy@{DIGEST}" - return "" - - builder = self._builder(sh) - builder.push_image("ci-dummy", "localhost:5000", None, "localhost:5000/ci-dummy") - - assert builder.pushed_digests["localhost:5000/ci-dummy"] == DIGEST - inspect_calls = [ - c for c in builder.console.sh.call_args_list if "docker image inspect" in c.args[0] - ] - assert len(inspect_calls) == 1 - assert "RepoDigests" in inspect_calls[0].args[0] - - def test_no_digest_anywhere_leaves_entry_absent_and_push_still_succeeds(self): - def sh(command, *args, **kwargs): - if "docker push" in command: - return "layer: Pushed" - if "docker image inspect" in command: - return "" - return "" - - builder = self._builder(sh) - result = builder.push_image("ci-dummy", "localhost:5000", None, "localhost:5000/ci-dummy") - - assert result == "localhost:5000/ci-dummy" - assert "localhost:5000/ci-dummy" not in builder.pushed_digests - # The gap is noted at dim level, not as a user-facing warning. - printed = " ".join(str(c) for c in builder.rich_console.print.call_args_list) - assert "[dim]" in printed - assert "no pushed digest" in printed.lower() - - def test_inspect_failure_is_swallowed(self): - def sh(command, *args, **kwargs): - if "docker push" in command: - return "layer: Pushed" - if "docker image inspect" in command: - raise RuntimeError("no such image") - return "" - - builder = self._builder(sh) - result = builder.push_image("ci-dummy", "localhost:5000", None, "localhost:5000/ci-dummy") - - assert result == "localhost:5000/ci-dummy" - assert builder.pushed_digests == {} - - def test_no_registry_records_nothing(self): - builder = self._builder(lambda command, *a, **k: "") - result = builder.push_image("ci-dummy") - - assert result == "ci-dummy" - assert builder.pushed_digests == {} -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_docker_builder.py -k PushImageRecordsDigest -v` -Expected: FAIL — `AttributeError: 'DockerBuilder' object has no attribute 'pushed_digests'` - -- [ ] **Step 3: Write the implementation** - -Add the import to `src/madengine/execution/docker_builder.py`, after the `from madengine.core.context import Context` line: - -```python -from madengine.core.image_digest import parse_push_digest, parse_repo_digest -``` - -In `DockerBuilder.__init__`, immediately after `self.built_images = {} # Track built images` (line 51), add: - -```python - self.pushed_digests = {} # registry_image -> digest recorded at push time -``` - -Replace the push block in `push_image()` (lines 394-399) — from `# Push the image` through `self.console.sh(push_command)` — with: - -```python - # Push the image - push_command = f"docker push {shlex.quote(registry_image)}" - self.rich_console.print(f"\n[bold blue]🚀 Starting docker push to registry...[/bold blue]") - print(f"📤 Registry: {registry}") - print(f"🏷️ Image: {registry_image}") - push_output = self.console.sh(push_command) - - self._record_pushed_digest(registry_image, push_output) -``` - -Add the new method immediately after `push_image()` (i.e. after its `raise` at line 408, before `def export_build_manifest`): - -```python - def _record_pushed_digest(self, registry_image: str, push_output: str) -> None: - """Record the digest of the image just pushed, for digest-pinned runs. - - Best-effort by design: the push has already succeeded by the time this - runs, so a missing digest is a manifest-completeness gap (noted at dim - level), never a build failure. Runs only consult the recorded digest - when --require-pinned-image is set. - - Args: - registry_image: The reference that was pushed. - push_output: stdout/stderr captured from the push command. - """ - digest = parse_push_digest(push_output) - - if not digest: - # Some registries/mirrors do not print the digest line; ask the - # daemon for the RepoDigests entry it recorded for this push. - try: - inspect_output = self.console.sh( - "docker image inspect --format '{{index .RepoDigests 0}}' " - + shlex.quote(registry_image) - ) - digest = parse_repo_digest(inspect_output) - except Exception: - digest = None - - if not digest: - self.rich_console.print( - f"[dim]No pushed digest recorded for {registry_image}; " - f"--require-pinned-image runs will reject this manifest entry[/dim]" - ) - return - - self.pushed_digests[registry_image] = digest - self.rich_console.print(f"[dim]Pushed digest: {registry_image} -> {digest}[/dim]") -``` - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_docker_builder.py -v` -Expected: PASS (all, including the 2 pre-existing naming tests) - -- [ ] **Step 5: Verify no existing push tests regressed** - -Run: `pytest tests/integration/test_docker_integration.py -k push_image -v` -Expected: PASS (6 tests). These mock `Console.sh` with `return_value="Success"`, so the digest parse returns None, the inspect fallback also returns `"Success"` (no digest), and `push_image` still returns the tag unchanged. - -- [ ] **Step 6: Commit** - -```bash -git add src/madengine/execution/docker_builder.py tests/unit/test_docker_builder.py -git commit -m "feat(build): capture pushed image digest during docker push" -``` - ---- - -## Task 4: Write `image_digest` into `build_info` at both push sites - -**Files:** -- Modify: `src/madengine/execution/docker_builder.py:762-772`, `:971-980` -- Test: `tests/unit/test_docker_builder.py` - -`build_info` is serialized wholesale into `build_manifest.json`, so setting the key here is all that is needed to get it into the manifest. - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_docker_builder.py`: - -```python -class TestBuildInfoCarriesImageDigest: - """Both push call sites copy the recorded digest into build_info.""" - - def _builder(self): - ctx = MagicMock() - ctx.ctx = {} - builder = DockerBuilder(ctx, MagicMock()) - builder.rich_console = MagicMock() - return builder - - def _run_single_arch(self, builder): - """Drive _build_model_single_arch with everything below push_image stubbed.""" - return builder._build_model_single_arch( - model_info={"name": "dummy", "dockerfile": "docker/dummy"}, - credentials={}, - clean_cache=False, - registry="localhost:5000", - phase_suffix="", - batch_build_metadata=None, - ) - - def test_single_arch_push_sets_image_digest(self): - builder = self._builder() - - def fake_push(docker_image, registry, credentials, explicit_registry_image): - builder.pushed_digests[explicit_registry_image] = DIGEST - return explicit_registry_image - - with patch.object( - builder, "_get_dockerfiles_for_model", return_value=["docker/dummy.ubuntu"] - ), patch.object( - builder, "build_image", return_value={"docker_image": "ci-dummy", "model": "dummy"} - ), patch.object( - builder, "_get_effective_gpu_architecture", return_value="" - ), patch.object( - builder, "_create_registry_image_name", return_value="localhost:5000/ci-dummy" - ), patch.object( - builder, "push_image", side_effect=fake_push - ): - results = self._run_single_arch(builder) - - assert results[0]["registry_image"] == "localhost:5000/ci-dummy" - assert results[0]["image_digest"] == DIGEST - # A push failure must still be recorded the way it is today. - assert "push_error" not in results[0] - - def test_single_arch_push_without_digest_omits_key(self): - builder = self._builder() - - with patch.object( - builder, "_get_dockerfiles_for_model", return_value=["docker/dummy.ubuntu"] - ), patch.object( - builder, "build_image", return_value={"docker_image": "ci-dummy", "model": "dummy"} - ), patch.object( - builder, "_get_effective_gpu_architecture", return_value="" - ), patch.object( - builder, "_create_registry_image_name", return_value="localhost:5000/ci-dummy" - ), patch.object( - builder, "push_image", return_value="localhost:5000/ci-dummy" - ): - results = self._run_single_arch(builder) - - assert results[0]["registry_image"] == "localhost:5000/ci-dummy" - assert "image_digest" not in results[0] -``` - -Extend the import at the top of `tests/unit/test_docker_builder.py`: - -```python -from unittest.mock import MagicMock, patch -``` - -> `_build_model_single_arch(self, model_info, credentials, clean_cache, registry, phase_suffix, batch_build_metadata)` is at `docker_builder.py:729`; the per-arch sibling is `_build_model_for_arch`, which takes an extra `arch` argument. Both are reached from `build_all_models` (`docker_builder.py:581`, `:617`). - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_docker_builder.py -k BuildInfoCarriesImageDigest -v` -Expected: FAIL — `KeyError: 'image_digest'` - -- [ ] **Step 3: Write the implementation** - -At the single-arch site (`docker_builder.py:769-770`), replace: - -```python - self.push_image(build_info["docker_image"], registry, credentials, registry_image) - build_info["registry_image"] = registry_image -``` - -with: - -```python - self.push_image(build_info["docker_image"], registry, credentials, registry_image) - build_info["registry_image"] = registry_image - # Recorded at push time; consumed only by --require-pinned-image runs. - pushed_digest = self.pushed_digests.get(registry_image) - if pushed_digest: - build_info["image_digest"] = pushed_digest -``` - -At the per-arch site (`docker_builder.py:977-978`), replace: - -```python - self.push_image(arch_image_name, registry, credentials, registry_image) - build_info["registry_image"] = registry_image -``` - -with: - -```python - self.push_image(arch_image_name, registry, credentials, registry_image) - build_info["registry_image"] = registry_image - # Recorded at push time; consumed only by --require-pinned-image runs. - pushed_digest = self.pushed_digests.get(registry_image) - if pushed_digest: - build_info["image_digest"] = pushed_digest -``` - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_docker_builder.py -v` -Expected: PASS - -- [ ] **Step 5: Confirm the manifest structure is unchanged for existing consumers** - -Run: `pytest tests/unit/test_orchestration.py tests/unit/test_slurm_multi.py -v` -Expected: PASS — `image_digest` is purely additive. - -- [ ] **Step 6: Commit** - -```bash -git add src/madengine/execution/docker_builder.py tests/unit/test_docker_builder.py -git commit -m "feat(build): record image_digest in build manifest entries" -``` - ---- - -## Task 5: `--require-pinned-image` flag and context propagation - -**Files:** -- Modify: `src/madengine/cli/commands/run.py:162-168` (flag), `:230-247` and `:318-341` (both `create_args_namespace` calls) -- Modify: `src/madengine/orchestration/run_orchestrator.py:85` (context merge), `:502` (manifest merge keys) -- Test: `tests/unit/test_orchestration.py` - -Two entry points must both work: the CLI flag, and the `require_pinned_image` key inside `--additional-context` (which is how CI pipelines drive madengine). Persisting the key into `manifest["context"]` is what carries the setting to the nested `madengine run` that the SLURM job script executes on each compute node. - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_orchestration.py`: - -```python -class TestRequirePinnedImageContext: - """--require-pinned-image and require_pinned_image both reach additional_context.""" - - @patch("madengine.orchestration.run_orchestrator.Context") - def test_cli_flag_sets_context_key(self, mock_context): - args = create_args_namespace( - additional_context=None, - require_pinned_image=True, - live_output=False, - ) - orch = RunOrchestrator(args) - assert orch.additional_context["require_pinned_image"] is True - - @patch("madengine.orchestration.run_orchestrator.Context") - def test_flag_absent_leaves_key_unset(self, mock_context): - args = create_args_namespace( - additional_context=None, - require_pinned_image=False, - live_output=False, - ) - orch = RunOrchestrator(args) - assert "require_pinned_image" not in orch.additional_context - - @patch("madengine.orchestration.run_orchestrator.Context") - def test_additional_context_key_alone_is_honoured(self, mock_context): - args = create_args_namespace( - additional_context="{'require_pinned_image': True}", - live_output=False, - ) - orch = RunOrchestrator(args) - assert orch.additional_context["require_pinned_image"] is True - - @patch("madengine.orchestration.run_orchestrator.Context") - def test_key_is_persisted_into_manifest_context(self, mock_context, tmp_path): - manifest_path = tmp_path / "build_manifest.json" - manifest_path.write_text(json.dumps({ - "built_images": {"img1": {"registry_image": "myorg/ci:m"}}, - "built_models": {"img1": {"name": "m"}}, - "context": {}, - "deployment_config": {}, - })) - - args = create_args_namespace( - additional_context=None, - require_pinned_image=True, - live_output=False, - ) - orch = RunOrchestrator(args) - orch._load_and_merge_manifest(str(manifest_path)) - - written = json.loads(manifest_path.read_text()) - assert written["context"]["require_pinned_image"] is True -``` - -Make sure `json`, `patch`, `create_args_namespace` and `RunOrchestrator` are imported at the top of `tests/unit/test_orchestration.py` — the existing `TestRunOrchestratorInit` and `TestCreateManifestFromLocalImage` classes already import most of these; add only what is missing: - -```python -from madengine.cli.utils import create_args_namespace -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_orchestration.py -k RequirePinnedImage -v` -Expected: FAIL — `KeyError: 'require_pinned_image'` - -- [ ] **Step 3: Write the implementation — orchestrator** - -In `src/madengine/orchestration/run_orchestrator.py`, immediately after `self.additional_context = merged_context` (line 85), add: - -```python - # The CLI flag and the require_pinned_image context key are equivalent; - # the key lets CI pipelines that drive madengine through - # --additional-context opt in the same way as for k8s/slurm/tools. - if getattr(args, "require_pinned_image", False): - self.additional_context["require_pinned_image"] = True -``` - -In `_load_and_merge_manifest`, extend the merge key list (line 502) from: - -```python - merge_keys = ["tools", "pre_scripts", "post_scripts", "encapsulate_script"] -``` - -to: - -```python - merge_keys = [ - "tools", - "pre_scripts", - "post_scripts", - "encapsulate_script", - # Persisted so nested runs on SLURM compute nodes (which re-enter - # `madengine run --manifest-file`) inherit the enforcement setting. - "require_pinned_image", - ] -``` - -- [ ] **Step 4: Write the implementation — CLI** - -In `src/madengine/cli/commands/run.py`, add a new option immediately after the `skip_model_run` option block (which ends at line 108 with `] = False,`): - -```python - require_pinned_image: Annotated[ - bool, - typer.Option( - "--require-pinned-image", - help=( - "Pull registry images by the digest recorded in the build manifest " - "instead of by tag. Fails immediately if the manifest has no digest " - "for an image. Equivalent to the 'require_pinned_image' " - "additional-context key." - ), - ), - ] = False, -``` - -Then add `require_pinned_image=require_pinned_image,` to **both** `create_args_namespace(...)` calls — in the manifest-exists branch (next to `skip_model_run=skip_model_run,` around line 245) and in the full-workflow branch (next to `skip_model_run=skip_model_run,` around line 339). - -- [ ] **Step 5: Run tests to verify they pass** - -Run: `pytest tests/unit/test_orchestration.py -v` -Expected: PASS - -- [ ] **Step 6: Verify the flag is wired into the CLI** - -Run: `madengine run --help | grep -A 3 require-pinned-image` -Expected: the option and its help text are listed. - -- [ ] **Step 7: Commit** - -```bash -git add src/madengine/cli/commands/run.py src/madengine/orchestration/run_orchestrator.py tests/unit/test_orchestration.py -git commit -m "feat(run): add --require-pinned-image flag and context propagation" -``` - ---- - -## Task 6: Enforce pinned pulls in local Docker execution - -**Files:** -- Modify: `src/madengine/execution/container_runner.py:2851-2859` -- Test: `tests/unit/test_container_runner.py` - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_container_runner.py`: - -```python -DIGEST = "sha256:" + "df36ef7e" * 8 - - -class TestRequirePinnedImageLocalRun: - """run_models_from_manifest honours require_pinned_image for registry pulls.""" - - def _manifest(self, tmpdir, build_info): - manifest_path = os.path.join(tmpdir, "build_manifest.json") - with open(manifest_path, "w") as f: - json.dump( - { - "built_images": {"img1": build_info}, - "built_models": { - "img1": {"name": "m", "tags": "t", "n_gpus": "1", "args": ""} - }, - }, - f, - ) - return manifest_path - - def _runner(self): - ctx = MagicMock() - ctx.ctx = {"docker_env_vars": {"MAD_SYSTEM_GPU_ARCHITECTURE": "gfx90a"}} - ctx.ensure_runtime_context = MagicMock() - console = MagicMock() - console.sh.return_value = "testhost" - runner = ContainerRunner(context=ctx, console=console) - runner.set_credentials({}) - return runner - - @patch("madengine.execution.container_runner.update_perf_csv") - def test_default_pulls_by_tag_even_when_digest_present(self, _mock_csv): - with tempfile.TemporaryDirectory() as tmpdir: - manifest_path = self._manifest( - tmpdir, - {"registry_image": "myorg/ci:m", "image_digest": DIGEST}, - ) - runner = self._runner() - runner.perf_csv_path = os.path.join(tmpdir, "perf.csv") - - with patch.object(runner, "pull_image") as mock_pull, patch.object( - runner, "run_container", return_value={"status": "SUCCESS"} - ): - runner.run_models_from_manifest(manifest_file=manifest_path, timeout=60) - - mock_pull.assert_called_once_with("myorg/ci:m") - - @patch("madengine.execution.container_runner.update_perf_csv") - def test_enabled_pulls_pinned_reference(self, _mock_csv): - with tempfile.TemporaryDirectory() as tmpdir: - manifest_path = self._manifest( - tmpdir, - {"registry_image": "myorg/ci:m", "image_digest": DIGEST}, - ) - runner = self._runner() - runner.perf_csv_path = os.path.join(tmpdir, "perf.csv") - runner.additional_context = {"require_pinned_image": True} - - with patch.object(runner, "pull_image") as mock_pull, patch.object( - runner, "run_container", return_value={"status": "SUCCESS"} - ) as mock_run: - runner.run_models_from_manifest(manifest_file=manifest_path, timeout=60) - - mock_pull.assert_called_once_with(f"myorg/ci@{DIGEST}") - # The container must run the same pinned reference that was pulled. - assert mock_run.call_args[1]["docker_image"] == f"myorg/ci@{DIGEST}" - - @patch("madengine.execution.container_runner.update_perf_csv") - def test_enabled_without_digest_fails_before_pulling(self, _mock_csv): - with tempfile.TemporaryDirectory() as tmpdir: - manifest_path = self._manifest(tmpdir, {"registry_image": "myorg/ci:m"}) - runner = self._runner() - runner.perf_csv_path = os.path.join(tmpdir, "perf.csv") - runner.additional_context = {"require_pinned_image": True} - - with patch.object(runner, "pull_image") as mock_pull, patch.object( - runner, "run_container" - ) as mock_run: - result = runner.run_models_from_manifest( - manifest_file=manifest_path, timeout=60 - ) - - mock_pull.assert_not_called() - mock_run.assert_not_called() - assert len(result["failed_runs"]) == 1 - assert "require-pinned-image" in result["failed_runs"][0]["error"] - - @patch("madengine.execution.container_runner.update_perf_csv") - def test_manifest_context_key_enables_enforcement(self, _mock_csv): - """A nested run on a SLURM compute node inherits the setting via manifest context.""" - with tempfile.TemporaryDirectory() as tmpdir: - manifest_path = os.path.join(tmpdir, "build_manifest.json") - with open(manifest_path, "w") as f: - json.dump( - { - "built_images": { - "img1": {"registry_image": "myorg/ci:m", "image_digest": DIGEST} - }, - "built_models": { - "img1": {"name": "m", "tags": "t", "n_gpus": "1", "args": ""} - }, - "context": {"require_pinned_image": True}, - }, - f, - ) - runner = self._runner() - runner.perf_csv_path = os.path.join(tmpdir, "perf.csv") - - with patch.object(runner, "pull_image") as mock_pull, patch.object( - runner, "run_container", return_value={"status": "SUCCESS"} - ): - runner.run_models_from_manifest(manifest_file=manifest_path, timeout=60) - - mock_pull.assert_called_once_with(f"myorg/ci@{DIGEST}") -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_container_runner.py -k RequirePinnedImageLocalRun -v` -Expected: FAIL — `test_enabled_pulls_pinned_reference` asserts the pinned ref but the tag is pulled. - -- [ ] **Step 3: Write the implementation** - -Add the import to `src/madengine/execution/container_runner.py`, after `from madengine.core.docker import Docker`: - -```python -from madengine.core.image_digest import resolve_pinned_image -``` - -Replace the registry branch at `container_runner.py:2851-2859`: - -```python - elif build_info.get("registry_image"): - # Registry image: Pull from registry - try: - self.pull_image(build_info["registry_image"]) - # Update docker_image to use registry image - run_image = build_info["registry_image"] - except Exception as pull_error: - self.rich_console.print(f"[yellow]Warning: Could not pull from registry, using local image[/yellow]") - run_image = image_name -``` - -with: - -```python - elif build_info.get("registry_image"): - # Registry image: Pull from registry. Under - # require_pinned_image this resolves to repo@sha256:... and - # raises (outside the pull try/except, so there is no tag - # fallback) when the manifest recorded no digest. - pull_target = resolve_pinned_image( - build_info["registry_image"], - build_info.get("image_digest"), - bool((self.additional_context or {}).get("require_pinned_image")), - model_name=model_info.get("name", ""), - ) - try: - self.pull_image(pull_target) - # Update docker_image to use registry image - run_image = pull_target - except Exception as pull_error: - self.rich_console.print(f"[yellow]Warning: Could not pull from registry, using local image[/yellow]") - run_image = image_name -``` - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_container_runner.py -v` -Expected: PASS - -- [ ] **Step 5: Commit** - -```bash -git add src/madengine/execution/container_runner.py tests/unit/test_container_runner.py -git commit -m "feat(run): pin local docker pulls to manifest digest when required" -``` - ---- - -## Task 7: Enforce pinned images in the Kubernetes pod spec - -**Files:** -- Modify: `src/madengine/deployment/k8s_template_context.py:518` -- Test: `tests/unit/test_k8s.py` - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_k8s.py`: - -```python -class TestK8sRequirePinnedImage: - """The generated pod spec image field honours require_pinned_image.""" - - DIGEST = "sha256:" + "df36ef7e" * 8 - - def _template_context(self, tmp_path, monkeypatch, require_pinned, image_digest): - """Build a real template context, the way prepare() does. - - _prepare_template_context reads the manifest and the model's scripts - directory from the current working directory, so the test runs inside - tmp_path with a minimal model tree. - """ - monkeypatch.chdir(tmp_path) - (tmp_path / "scripts" / "dummy").mkdir(parents=True) - (tmp_path / "scripts" / "dummy" / "run.sh").write_text("#!/bin/bash\necho hi\n") - - image_info = {"registry_image": "myorg/ci:m"} - if image_digest: - image_info["image_digest"] = image_digest - model_info = { - "name": "m", - "tags": ["t"], - "n_gpus": "1", - "args": "", - "scripts": "scripts/dummy/run.sh", - "dockerfile": "docker/dummy", - } - manifest = { - "built_images": {"img1": image_info}, - "built_models": {"img1": model_info}, - "context": {}, - } - (tmp_path / "build_manifest.json").write_text(json.dumps(manifest)) - - additional_context = { - "k8s": {"namespace": "default"}, - "gpu_vendor": "AMD", - "guest_os": "UBUNTU", - } - if require_pinned: - additional_context["require_pinned_image"] = True - - cfg = DeploymentConfig( - target="k8s", - manifest_file="build_manifest.json", - additional_context=additional_context, - ) - deployment = KubernetesDeployment(cfg) - return deployment._prepare_template_context(model_info, image_info) - - def test_default_uses_tag(self, tmp_path, monkeypatch): - ctx = self._template_context( - tmp_path, monkeypatch, require_pinned=False, image_digest=self.DIGEST - ) - assert ctx["image"] == "myorg/ci:m" - - def test_enabled_uses_pinned_reference(self, tmp_path, monkeypatch): - ctx = self._template_context( - tmp_path, monkeypatch, require_pinned=True, image_digest=self.DIGEST - ) - assert ctx["image"] == f"myorg/ci@{self.DIGEST}" - - def test_enabled_without_digest_raises(self, tmp_path, monkeypatch): - with pytest.raises(ConfigurationError): - self._template_context( - tmp_path, monkeypatch, require_pinned=True, image_digest=None - ) -``` - -Add whatever of these imports `tests/unit/test_k8s.py` is missing at the top (`pytest` is already there): - -```python -import json - -from madengine.core.errors import ConfigurationError -from madengine.deployment.base import DeploymentConfig -from madengine.deployment.kubernetes import KubernetesDeployment -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_k8s.py -k K8sRequirePinnedImage -v` -Expected: FAIL — `test_template_context_wires_the_resolver` fails; the resolver-behaviour tests pass already (they exercise Task 2 code directly, which is intentional: they document the K8s-facing contract). - -- [ ] **Step 3: Write the implementation** - -Add the import to `src/madengine/deployment/k8s_template_context.py`, next to the existing `from madengine.core.errors import ConfigurationError` (line 33): - -```python -from madengine.core.image_digest import resolve_pinned_image -``` - -In `_prepare_template_context`, immediately before the `return {` statement that begins the context dict, add: - -```python - # Under require_pinned_image the pod pulls repo@sha256:... so a moved tag - # surfaces as an ImagePullBackOff rather than a silent wrong-image run. - resolved_image = resolve_pinned_image( - image_info["registry_image"], - image_info.get("image_digest"), - bool(additional_context.get("require_pinned_image")), - model_name=model_name, - ) -``` - -Then change the image entry (line 518) from: - -```python - "image": image_info["registry_image"], -``` - -to: - -```python - "image": resolved_image, -``` - -> `additional_context` and `model_name` are both already local variables in this method (`additional_context = self.config.additional_context.copy()` and `model_name = model_info["name"]` near the top). - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_k8s.py -v` -Expected: PASS - -- [ ] **Step 5: Commit** - -```bash -git add src/madengine/deployment/k8s_template_context.py tests/unit/test_k8s.py -git commit -m "feat(k8s): pin pod image to manifest digest when required" -``` - ---- - -## Task 8: Enforce pinned images in the SLURM slurm_multi wrapper - -**Files:** -- Modify: `src/madengine/deployment/slurm.py:440-457` -- Test: `tests/unit/test_slurm_multi.py` - -The standard SLURM template path needs no change — it re-enters `madengine run --manifest-file` on each compute node, which goes through Task 6's local-Docker enforcement using the `require_pinned_image` key that Task 5 persisted into `manifest["context"]`. Only the self-managed `slurm_multi` wrapper, which bypasses that nested run, needs its own resolution. - -Setting `DOCKER_IMAGE_NAME` to the pinned reference pins both the parallel `srun docker pull` (which interpolates this variable) and the `docker run` inside the model's own script. See "Deliberate deviations" above. - -- [ ] **Step 1: Write the failing tests** - -Append to `tests/unit/test_slurm_multi.py`: - -```python -DIGEST = "sha256:" + "df36ef7e" * 8 - - -class TestSlurmMultiRequirePinnedImage: - """slurm_multi wrapper pins DOCKER_IMAGE_NAME (and thus the pull) when required.""" - - IMAGE_KEY = "rocm/pytorch-private:sglang_disagg_mori_20260502" - - def _deployment(self, tmp_path, require_pinned, image_digest): - script_rel = PR186_MODEL_ENTRY["scripts"] - script_abs = tmp_path / script_rel - script_abs.parent.mkdir(parents=True, exist_ok=True) - script_abs.write_text("#!/bin/bash\n# placeholder\n") - - image_entry = { - "image_name": self.IMAGE_KEY, - "docker_image": self.IMAGE_KEY, - "registry_image": self.IMAGE_KEY, - } - if image_digest: - image_entry["image_digest"] = image_digest - - manifest = { - "built_images": {self.IMAGE_KEY: image_entry}, - "built_models": {self.IMAGE_KEY: PR186_MODEL_ENTRY}, - "context": { - "docker_env_vars": {}, - "docker_mounts": {}, - "docker_build_arg": {}, - "gpu_vendor": "AMD", - "guest_os": "UBUNTU", - "docker_gpus": "all", - }, - } - manifest_path = tmp_path / "build_manifest.json" - manifest_path.write_text(json.dumps(manifest)) - - additional_context = { - "deploy": "slurm", - "gpu_vendor": "AMD", - "guest_os": "UBUNTU", - "slurm": dict( - PR186_MODEL_ENTRY["slurm"], output_dir=str(tmp_path / "slurm_results") - ), - "distributed": PR186_MODEL_ENTRY["distributed"], - } - if require_pinned: - additional_context["require_pinned_image"] = True - - cfg = DeploymentConfig( - target="slurm", - manifest_file=str(manifest_path), - additional_context=additional_context, - ) - return SlurmDeployment(cfg) - - def test_default_exports_tag(self, tmp_path): - dep = self._deployment(tmp_path, require_pinned=False, image_digest=DIGEST) - assert dep.prepare() is True - script_text = Path(dep.script_path).read_text() - assert f"export DOCKER_IMAGE_NAME={shlex.quote(self.IMAGE_KEY)}" in script_text - assert DIGEST not in script_text - - def test_enabled_exports_pinned_reference(self, tmp_path): - dep = self._deployment(tmp_path, require_pinned=True, image_digest=DIGEST) - assert dep.prepare() is True - script_text = Path(dep.script_path).read_text() - - pinned = f"rocm/pytorch-private@{DIGEST}" - assert f"export DOCKER_IMAGE_NAME={shlex.quote(pinned)}" in script_text - # The parallel pull interpolates the same value, so it is pinned too. - assert f"docker pull {pinned}" in script_text - - def test_enabled_without_digest_does_not_silently_fall_through(self, tmp_path): - """A missing digest must abort, not quietly take the standard template path. - - prepare()'s launcher peek wraps the slurm_multi dispatch in a bare - `except Exception: pass`. Without the re-raise added in Step 3b, a - ConfigurationError here would be swallowed and prepare() would generate - an ordinary (unpinned) sbatch script instead — the exact silent - degradation the flag exists to prevent. - """ - dep = self._deployment(tmp_path, require_pinned=True, image_digest=None) - with pytest.raises(ConfigurationError): - dep.prepare() -``` - -Add whatever of these imports `tests/unit/test_slurm_multi.py` is missing at the top: - -```python -import pytest - -from madengine.core.errors import ConfigurationError -``` - -- [ ] **Step 2: Run tests to verify they fail** - -Run: `pytest tests/unit/test_slurm_multi.py -k RequirePinnedImage -v` -Expected: FAIL — `test_enabled_exports_pinned_reference` finds the tag, not the pinned reference. - -- [ ] **Step 3: Write the implementation** - -Add the import to `src/madengine/deployment/slurm.py`, next to the other `madengine.*` imports (after `from madengine.utils.gpu_config import resolve_runtime_gpus`): - -```python -from madengine.core.image_digest import resolve_pinned_image -``` - -In `_prepare_slurm_multi_script`, replace the `DOCKER_IMAGE_NAME` resolution block (lines 440-457): - -```python - # Override DOCKER_IMAGE_NAME with the built image from manifest - # This ensures the run uses the freshly built image, not the base image - # Priority: docker_image_name param > model_info.docker_image > env_vars.DOCKER_IMAGE_NAME - if docker_image_name and docker_image_name.startswith("ci-"): - # The manifest key IS the built image name for madengine-built images - self.console.print(f"[cyan]Using built Docker image: {docker_image_name}[/cyan]") - env_vars["DOCKER_IMAGE_NAME"] = docker_image_name - elif "docker_image" in model_info: - built_image = model_info["docker_image"] - self.console.print(f"[cyan]Using Docker image: {built_image}[/cyan]") - env_vars["DOCKER_IMAGE_NAME"] = built_image - elif "image" in model_info: - # Fallback to 'image' field - built_image = model_info["image"] - self.console.print(f"[cyan]Using Docker image: {built_image}[/cyan]") - env_vars["DOCKER_IMAGE_NAME"] = built_image -``` - -with: - -```python - # Override DOCKER_IMAGE_NAME with the built image from manifest - # This ensures the run uses the freshly built image, not the base image - # Priority: docker_image_name param > model_info.docker_image > env_vars.DOCKER_IMAGE_NAME - if docker_image_name and docker_image_name.startswith("ci-"): - # The manifest key IS the built image name for madengine-built images - self.console.print(f"[cyan]Using built Docker image: {docker_image_name}[/cyan]") - env_vars["DOCKER_IMAGE_NAME"] = docker_image_name - elif "docker_image" in model_info: - built_image = model_info["docker_image"] - self.console.print(f"[cyan]Using Docker image: {built_image}[/cyan]") - env_vars["DOCKER_IMAGE_NAME"] = built_image - elif "image" in model_info: - # Fallback to 'image' field - built_image = model_info["image"] - self.console.print(f"[cyan]Using Docker image: {built_image}[/cyan]") - env_vars["DOCKER_IMAGE_NAME"] = built_image - - # Under require_pinned_image, pin DOCKER_IMAGE_NAME to the digest recorded - # at build time. slurm_multi runs the model's own script (no nested - # `madengine run` on the compute nodes), so enforcement has to happen here. - # Pinning the variable covers both the parallel `srun docker pull` below, - # which interpolates it, and the `docker run` inside the model script. - require_pinned = bool( - self.config.additional_context.get("require_pinned_image") - ) - if require_pinned and env_vars.get("DOCKER_IMAGE_NAME"): - image_entry = (self.manifest.get("built_images") or {}).get( - docker_image_name, {} - ) - env_vars["DOCKER_IMAGE_NAME"] = resolve_pinned_image( - env_vars["DOCKER_IMAGE_NAME"], - image_entry.get("image_digest"), - True, - model_name=model_info.get("name", ""), - ) - self.console.print( - f"[cyan]Pinned Docker image: {env_vars['DOCKER_IMAGE_NAME']}[/cyan]" - ) -``` - -- [ ] **Step 3b: Stop `prepare()` from swallowing the enforcement error** - -`prepare()` (`slurm.py:315-341`) wraps the whole slurm_multi dispatch — including the `_prepare_slurm_multi_script` call itself — in `except Exception: pass`, then falls through to the standard template path. Left as-is, a `ConfigurationError` from Step 3 would be silently discarded and an ordinary *unpinned* sbatch script would be generated instead. - -Narrow the handler so enforcement errors propagate. Replace the `except` clause at `slurm.py:339-341`: - -```python - except Exception: - # Fall through to develop's standard flow on any peek error - pass -``` - -with: - -```python - except ConfigurationError: - # Enforcement failures (e.g. --require-pinned-image with no recorded - # digest) are deliberate aborts, not peek errors. Falling through to - # the standard path here would silently generate an unpinned script. - raise - except Exception: - # Fall through to develop's standard flow on any peek error - pass -``` - -Add the import alongside the other `madengine.*` imports in `slurm.py`: - -```python -from madengine.core.errors import ConfigurationError -``` - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `pytest tests/unit/test_slurm_multi.py -v` -Expected: PASS - -- [ ] **Step 5: Commit** - -```bash -git add src/madengine/deployment/slurm.py tests/unit/test_slurm_multi.py -git commit -m "feat(slurm): pin slurm_multi image to manifest digest when required" -``` - ---- - -## Task 9: Lock in log-filename compatibility for pinned references - -**Files:** -- Test: `tests/unit/test_execution.py` - -`_docker_image_ref_for_log_naming()` already strips `@sha256:...`; this test prevents a future refactor from breaking it and silently changing log/tar filenames when pinning is on. - -- [ ] **Step 1: Write the test** - -Append to the existing `_docker_image_ref_for_log_naming` test class in `tests/unit/test_execution.py`: - -```python - def test_pinned_reference_names_same_as_untagged_reference(self): - digest = "sha256:" + "df36ef7e" * 8 - assert ( - _docker_image_ref_for_log_naming(f"registry/ns/myimg@{digest}") - == _docker_image_ref_for_log_naming("registry/ns/myimg") - ) - - def test_pinned_ci_reference_still_yields_tag(self): - digest = "sha256:" + "df36ef7e" * 8 - assert ( - _docker_image_ref_for_log_naming(f"rocm/ns/img:ci-m_model_df@{digest}") - == "ci-m_model_df" - ) -``` - -- [ ] **Step 2: Run the test** - -Run: `pytest tests/unit/test_execution.py -k log_naming -v` -Expected: PASS immediately — this is a characterization test of behaviour that already exists. If either assertion fails, stop and report it; that would mean pinned references change log filenames and the spec's compatibility claim is wrong. - -- [ ] **Step 3: Commit** - -```bash -git add tests/unit/test_execution.py -git commit -m "test(execution): cover log naming for digest-pinned image references" -``` - ---- - -## Task 10: Documentation and full-suite verification - -**Files:** -- Modify: `docs/cli-reference.md:232` (run options table) -- Modify: `docs/configuration.md` - -- [ ] **Step 1: Add the CLI reference row** - -In `docs/cli-reference.md`, insert a new row into the `run` options table immediately after the `--skip-model-run` row (line 232): - -```markdown -| `--require-pinned-image` | | FLAG | `False` | Pull registry images by the `sha256` digest recorded in the build manifest (`repo@sha256:...`) instead of by tag, so a tag that moved between build and run fails loudly instead of silently running a different image. Fails immediately — with no tag fallback — if the manifest has no digest for an image. Equivalent to the `require_pinned_image` additional-context key. See [Configuration — Pinned image digests](configuration.md#pinned-image-digests). | -``` - -- [ ] **Step 2: Add the configuration section** - -In `docs/configuration.md`, add a new section after the "Run phase: log error pattern scan" section (which ends before "## System environment collection (rocEnvTool)"): - -```markdown -## Pinned image digests - -Every build records the digest of the image it pushes as `image_digest` on each -`built_images` entry in `build_manifest.json`. This capture is always on and -costs nothing: by default the digest is carried along and never used. - -Pass `--require-pinned-image` (or set `"require_pinned_image": true` in -`--additional-context`) to make the run phase pull `repo@sha256:...` instead of -the tag: - -```bash -madengine run --manifest-file build_manifest.json --require-pinned-image - -# Equivalent, for pipelines that drive madengine through additional context -madengine run --manifest-file build_manifest.json \ - --additional-context "{'require_pinned_image': True}" -``` - -| Behaviour | Flag absent (default) | Flag set | -|---|---|---| -| Registry pull | By tag | By digest (`repo@sha256:...`) | -| Manifest has no `image_digest` | Pull by tag | **Fails immediately**, no tag fallback | -| Tag moved since the build | Silently runs the newer image | Registry rejects the pull (`manifest unknown`) | - -Applies to all three execution paths: local Docker, Kubernetes (the pod spec -`image` field), and SLURM. On SLURM the setting is written into the manifest's -`context` block so the nested `madengine run` on each compute node inherits it. - -**Limitations** - -- This does not prevent two concurrent builds from racing to push the same - mutable tag. It converts the resulting silent wrong-image run into a fast, - clear failure. Eliminating the race requires unique tags per build in the - calling CI pipeline. -- Manifests produced by the build-on-compute-node path (SLURM batch builds, - which push from inside a generated sbatch script) carry no `image_digest`. - Runs against those manifests fail fast when the flag is set. -``` - -- [ ] **Step 3: Run the full unit suite** - -Run: `pytest tests/unit -v` -Expected: PASS, no regressions. - -- [ ] **Step 4: Run the integration suite** - -Run: `pytest tests/integration -v -m "not slow"` -Expected: PASS. Pay particular attention to `tests/integration/test_docker_integration.py -k push_image`, which asserts the exact `docker tag` / `docker push` call shapes. - -- [ ] **Step 5: Format and lint the changed files** - -```bash -black src/madengine/core/image_digest.py src/madengine/execution/docker_builder.py \ - src/madengine/execution/container_runner.py src/madengine/deployment/slurm.py \ - src/madengine/deployment/k8s_template_context.py \ - src/madengine/orchestration/run_orchestrator.py src/madengine/cli/commands/run.py \ - tests/unit/test_image_digest.py -isort src/madengine/core/image_digest.py tests/unit/test_image_digest.py -mypy src/madengine/core/image_digest.py -``` - -Expected: `black`/`isort` reformat or report no changes; `mypy` reports no errors in the new module. - -> Only run `black`/`isort` on the files you actually touched. Reformatting untouched files would bloat the diff. - -- [ ] **Step 6: Commit** - -```bash -git add docs/cli-reference.md docs/configuration.md -git commit -m "docs: document --require-pinned-image and image digest capture" -``` - -- [ ] **Step 7: Manual end-to-end smoke check (optional, requires a registry)** - -```bash -# Build and push, then confirm the digest landed in the manifest -madengine build --tags dummy --registry localhost:5000 -python -c "import json; m=json.load(open('build_manifest.json')); print({k: v.get('image_digest') for k, v in m['built_images'].items()})" - -# Run with enforcement and confirm the pull is by digest -madengine run --manifest-file build_manifest.json --require-pinned-image --live-output 2>&1 | grep "docker pull" -``` - -Expected: the manifest prints a `sha256:...` per image, and the pull line contains `@sha256:`. - ---- - -## Verification checklist against the spec's testing plan - -| Spec test | Covered by | -|---|---| -| 1. Push output digest → `image_digest` | Task 3 `test_digest_parsed_from_push_output` + Task 4 `test_single_arch_push_sets_image_digest` | -| 2. Fallback to `docker image inspect` | Task 3 `test_falls_back_to_image_inspect_when_push_output_has_no_digest` | -| 3. Both paths fail → absent key, debug log, build unaffected | Task 3 `test_no_digest_anywhere_leaves_entry_absent_and_push_still_succeeds`, `test_inspect_failure_is_swallowed` | -| 4. Flag absent → pull by tag, no new output | Task 6 `test_default_pulls_by_tag_even_when_digest_present`, Task 7 `test_default_uses_tag`, Task 8 `test_default_exports_tag` | -| 5. Flag present + digest → pinned reference | Task 6 `test_enabled_pulls_pinned_reference`, Task 7 `test_enabled_uses_pinned_reference`, Task 8 `test_enabled_exports_pinned_reference` | -| 6. Flag present, no digest → fail before pull | Task 2 `test_enabled_without_digest_raises`, Task 6 `test_enabled_without_digest_fails_before_pulling` | -| 7. K8s pod spec and SLURM script carry pinned/tag reference | Task 7 + Task 8 | -| 8. Log filename derivation unchanged | Task 9 | -| 9. Existing `docker_sha` / manifest tests unmodified | Task 4 Step 5, Task 10 Steps 3-4 | diff --git a/docs/superpowers/specs/2026-08-27-pinned-image-digest-design.md b/docs/superpowers/specs/2026-08-27-pinned-image-digest-design.md deleted file mode 100644 index 08d69aea..00000000 --- a/docs/superpowers/specs/2026-08-27-pinned-image-digest-design.md +++ /dev/null @@ -1,156 +0,0 @@ -# Design: pin registry pulls to the digest recorded at build time - -## Problem - -A client-perf-hub accuracy run failed because the image pulled at run time -did not match the image pushed at build time: - -- build pushed `sha256:df36ef7e...` -- the CI runner later pulled `sha256:cbb7e5ed...` - -Both shared the same registry tag. A concurrent build of the same model -pushed a newer image to that tag between the two events, so the run phase -silently executed a different image than the one the build phase produced. - -It was assumed madengine already guards against this ("we use it to confirm -if the pulled image is identical to the one in the build manifest"). It does -not. Tracing the code: - -- `build_info["docker_sha"]` (`docker_builder.py:303-308`) is the digest of - the Dockerfile's `FROM` (base) image, not the image madengine builds and - pushes. It is consumed only as a reporting column (`update_perf_csv.py`, - `k8s_results.py`, `slurm.py`). -- `push_image()` (`docker_builder.py:395`) discards `docker push` stdout, - which is where the pushed digest (`digest: sha256:...`) is printed. It is - never captured anywhere. -- The run phase pulls by tag only (`container_runner.py:2854` → - `pull_image()` at `container_runner.py:572`, a plain `docker pull `) - and performs no comparison against anything in the manifest. The only - image-identity checks in the codebase (`_local_image_id`, - `BUILD_FINGERPRINT_LABEL`, `container_runner.py:2464-2503`) exist solely - for cross-node consistency in multi-node SLURM *local-image* mode; they - never touch registry digests. - -This design adds the missing capability: capture the real pushed digest at -build time, and optionally enforce it at run time. - -## Non-goals - -- This does not prevent the tag race itself. Two builds pushing the same - mutable tag (`f"{registry}:{model_name}"`, `build_orchestrator.py:1055`) - will still clobber each other; the loser under the new flag fails fast - with a clear error instead of silently running the wrong image. Actually - eliminating the race (e.g. unique tags per build) is a client-perf-hub / - upstream CI change, out of scope here. -- This does not rename or repurpose `build_info["docker_sha"]`. It is wired - into four existing reporting sinks and column orders; changing its - meaning is a separate, riskier change and is called out only as a - possible future cleanup. -- This does not change default pull behavior for any existing user. See - "Opt-in enforcement" below. - -## Design - -### 1. Capture the pushed digest at build time (always on, default mode included) - -In `DockerBuilder.push_image()` (`docker_builder.py`), after `docker push` -succeeds, parse the digest from its output (`digest: sha256:...`, the same -line format already parsed for base-image SHA in `docker_builder.py:305`). - -If the push output doesn't contain a parseable digest line (registry output -format variance, e.g. some mirrors), fall back to: -``` -docker image inspect --format '{{index .RepoDigests 0}}' -``` -and extract the digest from that. - -If both fail to produce a digest, log a **debug/dim-level** note (not a -user-facing warning) and continue — this is a manifest-completeness gap to -leave a trail for later, not a failure. Push already succeeded; we don't -block on this. - -Store the result as a new field: `build_info["image_digest"]` (e.g. -`"sha256:df36ef7e..."`), separate from `docker_sha`. This flows into -`build_manifest.json` through the existing `built_images` / -`export_build_manifest()` path with no other changes needed. - -This capture step runs unconditionally — no flag gates it. It's the -recording half of the fix, and it's inert (never read by pull logic) unless -strict mode is on. - -### 2. Opt-in enforcement at run time - -New flag: `--require-pinned-image` on the `run` command, mirrored as an -`additional_context` key (e.g. `"require_pinned_image": true`) so CI -pipelines that drive madengine through `additional_context` rather than raw -CLI flags can set it the same way as other behavior-affecting keys -(`k8s`, `slurm`, `tools`, ...). - -**Default (flag absent):** no behavior change whatsoever. Pull by tag, -exactly as today. `image_digest` rides along in the manifest unused. - -**With the flag set**, for every `build_info` entry with a -`registry_image`: -- If `image_digest` is present, construct a pinned reference - `repo@sha256:...` (stripping any existing tag) and pull that image - instead of the tag. The registry itself now enforces identity: a moved - tag surfaces as a normal "manifest unknown" pull failure rather than a - silent wrong-image success. -- If `image_digest` is absent (older manifest, or capture failed at build - time), fail immediately with a clear error naming the model/image and - explaining that the manifest has no recorded digest — do not fall back to - pulling by tag. Silent degradation defeats the purpose of asking for the - guarantee. - -One shared helper (e.g. `build_pinned_reference(registry_image, digest)`) -builds the `repo@sha256:...` string, used identically by all three -enforcement call sites so they can't drift: - -| Path | Location | Change under the flag | -|---|---|---| -| Local Docker | `container_runner.py:2854` | `pull_image(pinned_ref)` | -| Kubernetes | `k8s_template_context.py:518` (`"image": image_info["registry_image"]`) | `"image": pinned_ref` | -| SLURM | `slurm.py:544` (parallel `srun` pull) | pinned ref substituted into the generated pull command | - -### 3. Compatibility check: log filenames - -`container_runner_helpers.py:256` already strips `@sha256:...` before -deriving log/tar filenames from an image reference (`ref_without_digest = -s.split("@", 1)[0]`). Pinned references flow through this unchanged — no -new filename collisions. Confirmed by reading the function; will be -covered by a test regardless. - -## Testing plan - -TDD; each behavior below is a separate test: - -1. `docker push` output containing a `digest: sha256:...` line → - `build_info["image_digest"]` set to that value. -2. `docker push` output without a parseable digest line → falls back to - `docker image inspect --format '{{index .RepoDigests 0}}'`. -3. Both parse paths fail → `image_digest` absent, a debug-level log line is - emitted, and the build/push otherwise succeeds unaffected. -4. Flag absent (default) → run pulls by tag regardless of whether - `image_digest` is present in the manifest; no new log output. -5. Flag present, `image_digest` present → the pull command (or k8s image - field / SLURM pull line) contains `repo@sha256:...`. -6. Flag present, `image_digest` absent → run fails immediately with a clear - error, before any pull is attempted. -7. K8s pod spec and SLURM generated script both carry the pinned reference - when the flag is set, and the untouched tag reference when it isn't. -8. Log/tar filename derivation given a pinned (`@sha256:...`) reference - produces the same filename as today's digest-free reference. -9. Existing tests touching `docker_sha` / `build_manifest.json` structure - continue to pass unmodified (the new field is additive). - -## Rollout note (for the reply to Tej/Rahul) - -- The verification described in the thread ("we use it to confirm...") - does not currently exist in the code; this design builds it, gated - behind `--require-pinned-image` / `require_pinned_image` context key. -- Turning the flag on stops the *symptom* (silently running the wrong - image) by failing fast instead. It does not stop the *cause* (two builds - racing to push the same mutable tag) — that needs a change on the - client-perf-hub / CI side (e.g. unique tags per build). -- client-perf-hub must explicitly opt in for this guarantee to apply to its - runs; it is not automatic on upgrade. diff --git a/src/madengine/deployment/config_loader.py b/src/madengine/deployment/config_loader.py index 06d8a1b1..4398b5f3 100644 --- a/src/madengine/deployment/config_loader.py +++ b/src/madengine/deployment/config_loader.py @@ -230,20 +230,23 @@ def infer_and_validate_deploy_type(cls, user_config: Dict[str, Any]) -> str: Infer deployment type from config structure and validate for conflicts. Convention over Configuration: Presence of k8s/slurm field indicates deployment intent. - + The SLURM flavor ("slurm" vs "spur") is distinguished by slurm.scheduler, + because spur ships SLURM-compatible CLI shims and reuses the same block. + Args: user_config: User configuration dictionary - + Returns: - Deployment type: "k8s", "slurm", or "local" - + Deployment type: "k8s", "spur", "slurm", or "local" + Raises: ValueError: If conflicting deployment configs present """ has_k8s = "k8s" in user_config or "kubernetes" in user_config has_slurm = "slurm" in user_config explicit_deploy = user_config.get("deploy", "").lower() - + scheduler = str((user_config.get("slurm") or {}).get("scheduler", "") or "").lower() + # Validation Rule 1: Can't have both k8s and slurm configs if has_k8s and has_slurm: raise ValueError( @@ -258,22 +261,34 @@ def infer_and_validate_deploy_type(cls, user_config: Dict[str, Any]) -> str: f"Conflicting deployment: 'deploy' field is '{explicit_deploy}' but no 'k8s' config present. " "Either add 'k8s' config or remove 'deploy' field." ) - if explicit_deploy == "slurm" and not has_slurm: + if explicit_deploy in ["slurm", "spur"] and not has_slurm: raise ValueError( - f"Conflicting deployment: 'deploy' field is 'slurm' but no 'slurm' config present. " + f"Conflicting deployment: 'deploy' field is '{explicit_deploy}' but no 'slurm' config present. " "Either add 'slurm' config or remove 'deploy' field." ) + if explicit_deploy == "spur" and scheduler not in ("", "spur"): + raise ValueError( + f"Conflicting deployment: 'deploy' field is 'spur' but slurm.scheduler is '{scheduler}'. " + "Set slurm.scheduler to 'spur' or remove 'deploy' field." + ) if explicit_deploy == "local" and (has_k8s or has_slurm): raise ValueError( f"Conflicting deployment: 'deploy' field is 'local' but k8s/slurm config present. " "Remove k8s/slurm config for local execution." ) - + + # Validation Rule 3: slurm.scheduler must name a known SLURM flavor + if has_slurm and scheduler not in ("", "slurm", "spur"): + raise ValueError( + f"Unknown slurm.scheduler '{scheduler}'. Supported values are 'slurm' (default) and 'spur'." + ) + # Infer deployment type from config presence if has_k8s: return "k8s" elif has_slurm: - return "slurm" + # Spur reuses the "slurm" block; the flavor comes from slurm.scheduler. + return "spur" if scheduler == "spur" or explicit_deploy == "spur" else "slurm" else: return "local" @@ -308,7 +323,8 @@ def load_config(cls, user_config: Dict[str, Any]) -> Dict[str, Any]: # Note: We do NOT add a "deploy" field - type is inferred from structure if deploy_type == "k8s": return cls.load_k8s_config(user_config) - elif deploy_type == "slurm": + elif deploy_type in ("slurm", "spur"): + # Spur reuses the SLURM presets; only the scheduler flavor differs. return cls.load_slurm_config(user_config) else: # Local - return as-is (no deploy field needed) diff --git a/src/madengine/deployment/factory.py b/src/madengine/deployment/factory.py index 944a4022..5bbba46a 100644 --- a/src/madengine/deployment/factory.py +++ b/src/madengine/deployment/factory.py @@ -81,6 +81,11 @@ def register_default_deployments(): DeploymentFactory.register("slurm", SlurmDeployment) + # Spur (Crusoe) scheduler: SLURM-compatible CLI with job-array fan-out. + from .spur import SpurDeployment + + DeploymentFactory.register("spur", SpurDeployment) + # Register Kubernetes if library is available try: from .kubernetes import KubernetesDeployment diff --git a/src/madengine/deployment/slurm.py b/src/madengine/deployment/slurm.py index a83ebcb1..ebc568c0 100644 --- a/src/madengine/deployment/slurm.py +++ b/src/madengine/deployment/slurm.py @@ -11,7 +11,9 @@ Copyright (c) Advanced Micro Devices, Inc. All rights reserved. """ +import json import os +import re import shlex import shutil import subprocess @@ -19,8 +21,20 @@ from pathlib import Path from typing import Any, Dict, List, Optional -from .base import BaseDeployment, DeploymentConfig, DeploymentResult, DeploymentStatus, create_jinja_env -from .primus_backend import infer_primus_backend_from_model_name, merged_primus_config +from madengine.core.errors import ConfigurationError +from madengine.core.image_digest import resolve_pinned_image +from madengine.core.timeout import subprocess_timeout +from madengine.utils.gpu_config import resolve_runtime_gpus +from madengine.utils.path_utils import scripts_base_dir_from +from madengine.utils.run_details import get_build_number, get_pipeline + +from .base import ( + BaseDeployment, + DeploymentConfig, + DeploymentResult, + DeploymentStatus, + create_jinja_env, +) from .common import ( canonicalize_distributed_launcher, configure_multi_node_profiling, @@ -28,14 +42,8 @@ normalize_launcher, ) from .config_loader import ConfigLoader, apply_deployment_config +from .primus_backend import infer_primus_backend_from_model_name, merged_primus_config from .slurm_node_selector import SlurmNodeSelector -from madengine.core.errors import ConfigurationError -from madengine.core.image_digest import resolve_pinned_image -from madengine.core.timeout import subprocess_timeout -from madengine.utils.gpu_config import resolve_runtime_gpus -from madengine.utils.run_details import get_build_number, get_pipeline -from madengine.utils.path_utils import scripts_base_dir_from -import json class SlurmDeployment(BaseDeployment): @@ -59,6 +67,10 @@ class SlurmDeployment(BaseDeployment): DEPLOYMENT_TYPE = "slurm" REQUIRED_TOOLS = ["sbatch", "squeue", "scontrol"] # Must be available locally + # Set by the spur subclass. When True, multi-node fan-out is achieved with a + # job array (one single-node task per node) instead of `srun`, because spur's + # srun cannot dispatch tasks to other nodes. See SpurDeployment. + IS_SPUR = False def __init__(self, config: DeploymentConfig): """ @@ -96,7 +108,7 @@ def __init__(self, config: DeploymentConfig): self.inside_allocation = os.environ.get("SLURM_JOB_ID") is not None self.existing_job_id = os.environ.get("SLURM_JOB_ID", "") self.allocation_nodes = self._get_allocation_node_count() - + if self.inside_allocation: self.console.print( f"[cyan]✓ Detected existing SLURM allocation: Job {self.existing_job_id}[/cyan]" @@ -105,18 +117,50 @@ def __init__(self, config: DeploymentConfig): f" Allocation has {self.allocation_nodes} nodes available" ) + @staticmethod + def _expand_nodelist(nodelist: str) -> List[str]: + """Expand a SLURM nodelist string into a list of hostnames. + + Stock SLURM emits a compressed form (e.g. "node[01-03,05]") that needs + `scontrol show hostnames` to expand. Spur (Crusoe) instead already + exposes an expanded comma-separated form (e.g. "nodeA,nodeB") in + SLURM_NODELIST and does NOT implement `scontrol show hostname[s]`. + + Strategy: if the string has no range/brace syntax, just split on commas + (works on spur). Otherwise try `scontrol show hostnames` and fall back to + a naive comma split if that is unavailable. + """ + if not nodelist: + return [] + nodelist = nodelist.strip() + if "[" not in nodelist: + return [h for h in re.split(r"[,\s]+", nodelist) if h] + try: + result = subprocess.run( + ["scontrol", "show", "hostnames", nodelist], + capture_output=True, + text=True, + timeout=10, + ) + if result.returncode == 0 and result.stdout.strip(): + return [h for h in result.stdout.split("\n") if h.strip()] + except Exception: + pass + # Best-effort fallback (cannot expand ranges without scontrol). + return [h for h in re.split(r"[,\s]+", nodelist) if h] + def _get_allocation_node_count(self) -> int: """ Get number of nodes in current SLURM allocation. - + Note: SLURM_NNODES reflects the current job step, not the full allocation. We query the job directly using scontrol to get the actual node count. """ if not self.inside_allocation: return 0 - + job_id = self.existing_job_id - + # Query the actual job's node count using scontrol (most accurate) try: result = subprocess.run( @@ -138,7 +182,7 @@ def _get_allocation_node_count(self) -> int: pass except Exception: pass - + # Fallback: Try SLURM_JOB_NUM_NODES (full job node count, if set) job_num_nodes = os.environ.get("SLURM_JOB_NUM_NODES") if job_num_nodes: @@ -146,7 +190,7 @@ def _get_allocation_node_count(self) -> int: return int(job_num_nodes) except ValueError: pass - + # Fallback: SLURM_NNODES (may be step-specific, not full allocation) nnodes = os.environ.get("SLURM_NNODES") if nnodes: @@ -154,59 +198,49 @@ def _get_allocation_node_count(self) -> int: return int(nnodes) except ValueError: pass - - # Last resort: count nodes in SLURM_NODELIST + + # Last resort: count nodes in SLURM_NODELIST (spur-safe expansion) nodelist = os.environ.get("SLURM_NODELIST") if nodelist: - try: - result = subprocess.run( - ["scontrol", "show", "hostname", nodelist], - capture_output=True, - text=True, - timeout=10, - ) - if result.returncode == 0: - return len(result.stdout.strip().split("\n")) - except Exception: - pass - + expanded = self._expand_nodelist(nodelist) + if expanded: + return len(expanded) + return 0 def _validate_allocation_nodes(self) -> tuple: """ Validate that existing allocation has enough nodes for the job. - + Returns: Tuple of (is_valid, error_message) """ if not self.inside_allocation: return True, "" - + requested_nodes = self.nodes available_nodes = self.allocation_nodes - + if available_nodes < requested_nodes: return False, ( f"Insufficient nodes in current allocation. " f"Requested: {requested_nodes}, Available: {available_nodes}. " f"Either reduce nodes in config or use a larger allocation." ) - + if available_nodes > requested_nodes: self.console.print( f"[yellow]⚠ Note: Using {requested_nodes} of {available_nodes} " f"available nodes in allocation[/yellow]" ) - + return True, "" def validate(self) -> bool: """Validate SLURM commands are available locally.""" # Check required SLURM CLI tools for tool in self.REQUIRED_TOOLS: - result = subprocess.run( - ["which", tool], capture_output=True, timeout=5 - ) + result = subprocess.run(["which", tool], capture_output=True, timeout=5) if result.returncode != 0: self.console.print( f"[red]✗ Required tool not found: {tool}[/red]\n" @@ -226,7 +260,9 @@ def validate(self) -> bool: return False if self.gpus_per_node < 1: - self.console.print(f"[red]✗ Invalid GPUs per node: {self.gpus_per_node}[/red]") + self.console.print( + f"[red]✗ Invalid GPUs per node: {self.gpus_per_node}[/red]" + ) return False self.console.print("[green]✓ SLURM environment validated[/green]") @@ -251,10 +287,10 @@ def _submission_bin_dir() -> Optional[str]: def _validate_cli_availability(self) -> bool: """ Validate madengine is available before job submission. - + Compute nodes inherit the submission environment, so madengine must be available in PATH on the submission node. - + Returns: bool: True if madengine is available and functional """ @@ -266,38 +302,31 @@ def _validate_cli_availability(self) -> bool: # A cold import off shared/NFS storage can take far longer than a # local one, so this only guards against a hung interpreter. timeout=600, - check=False + check=False, ) if result.returncode == 0: version = result.stdout.strip() or "unknown" self.console.print( f"[green]✓[/green] madengine available: [cyan]{version}[/cyan]" ) - + # Show path for transparency which_result = subprocess.run( - ["which", "madengine"], - capture_output=True, - text=True, - check=False + ["which", "madengine"], capture_output=True, text=True, check=False ) if which_result.returncode == 0: cli_path = which_result.stdout.strip() self.console.print(f" Path: [dim]{cli_path}[/dim]") - + return True else: - self.console.print( - "[red]✗ madengine found but returned error[/red]" - ) + self.console.print("[red]✗ madengine found but returned error[/red]") if result.stderr: self.console.print(f" Error: {result.stderr.strip()}") return False - + except FileNotFoundError: - self.console.print( - "\n[red]✗ ERROR: madengine not found[/red]\n" - ) + self.console.print("\n[red]✗ ERROR: madengine not found[/red]\n") self.console.print( "[yellow]Compute nodes need madengine in PATH.[/yellow]\n" "\n[bold]To fix:[/bold]\n" @@ -327,10 +356,9 @@ def prepare(self) -> bool: if model_keys_peek: model_info_peek = self.manifest["built_models"][model_keys_peek[0]] model_distributed_peek = model_info_peek.get("distributed", {}) - launcher_type_peek = ( - model_distributed_peek.get("launcher") - or self.distributed_config.get("launcher", "torchrun") - ) + launcher_type_peek = model_distributed_peek.get( + "launcher" + ) or self.distributed_config.get("launcher", "torchrun") if is_self_managed_launcher(launcher_type_peek): self.output_dir.mkdir(parents=True, exist_ok=True) self.console.print( @@ -355,7 +383,7 @@ def prepare(self) -> bool: "\n[yellow]⚠ Tip: Compute nodes inherit your submission environment[/yellow]" ) return False - + try: self.output_dir.mkdir(parents=True, exist_ok=True) @@ -395,7 +423,9 @@ def _normalize_nodelist(nodelist: Optional[str]) -> Optional[str]: return None return ",".join(n.strip() for n in nodelist.split(",") if n.strip()) - def _prepare_slurm_multi_script(self, model_info: Dict, docker_image_name: str = None) -> bool: + def _prepare_slurm_multi_script( + self, model_info: Dict, docker_image_name: str = None + ) -> bool: """ Escape hatch for self-orchestrating multi-container SLURM topologies. @@ -416,41 +446,47 @@ def _prepare_slurm_multi_script(self, model_info: Dict, docker_image_name: str = if not model_script: self.console.print("[red]✗ No scripts defined in model_info[/red]") return False - + # Get manifest directory (where the model script is relative to) manifest_dir = Path(self.config.manifest_file).parent.absolute() model_script_path = manifest_dir / model_script - + if not model_script_path.exists(): - self.console.print(f"[red]✗ Model script not found: {model_script_path}[/red]") + self.console.print( + f"[red]✗ Model script not found: {model_script_path}[/red]" + ) return False - + # Get environment variables env_vars = {} - + # From model_info.env_vars if "env_vars" in model_info: env_vars.update(model_info["env_vars"]) - + # From additional_context.env_vars if "env_vars" in self.config.additional_context: env_vars.update(self.config.additional_context["env_vars"]) - + # From distributed config (model's distributed section) model_distributed = model_info.get("distributed", {}) - sglang_disagg_config = model_distributed.get("sglang_disagg", {}) or self.distributed_config.get("sglang_disagg", {}) + sglang_disagg_config = model_distributed.get( + "sglang_disagg", {} + ) or self.distributed_config.get("sglang_disagg", {}) if sglang_disagg_config: if "xP" not in env_vars: env_vars["xP"] = str(sglang_disagg_config.get("prefill_nodes", 1)) if "yD" not in env_vars: env_vars["yD"] = str(sglang_disagg_config.get("decode_nodes", 1)) - + # Override DOCKER_IMAGE_NAME with the built image from manifest # This ensures the run uses the freshly built image, not the base image # Priority: docker_image_name param > model_info.docker_image > env_vars.DOCKER_IMAGE_NAME if docker_image_name and docker_image_name.startswith("ci-"): # The manifest key IS the built image name for madengine-built images - self.console.print(f"[cyan]Using built Docker image: {docker_image_name}[/cyan]") + self.console.print( + f"[cyan]Using built Docker image: {docker_image_name}[/cyan]" + ) env_vars["DOCKER_IMAGE_NAME"] = docker_image_name elif "docker_image" in model_info: built_image = model_info["docker_image"] @@ -496,141 +532,242 @@ def _prepare_slurm_multi_script(self, model_info: Dict, docker_image_name: str = else "" ) _bash_invocation = f"bash {_script_name_q} {_model_args_q}".rstrip() - + # Generate simple wrapper script # IMPORTANT: SBATCH directives MUST be at the top, right after #!/bin/bash script_lines = [ "#!/bin/bash", f"#SBATCH --job-name=madengine-{model_info['name']}", - f"#SBATCH --output={self.output_dir}/madengine-{model_info['name']}_%j_%t.out", - f"#SBATCH --error={self.output_dir}/madengine-{model_info['name']}_%j_%t.err", - f"#SBATCH --partition={self.partition}", - f"#SBATCH --nodes={self.nodes}", - f"#SBATCH --ntasks={self.nodes}", ] + if self.IS_SPUR: + # spur: one single-node task per node via a job array (srun cannot fan out). + # Log names MUST use %A (array job id, == the id sbatch returns) rather + # than %j (each array task's own job id), so collect_results() and + # _show_log_summary() can find them by deployment_id. %a is the node rank. + script_lines.extend( + [ + f"#SBATCH --output={self.output_dir}/madengine-{model_info['name']}_%A_%a.out", + f"#SBATCH --error={self.output_dir}/madengine-{model_info['name']}_%A_%a.err", + f"#SBATCH --partition={self.partition}", + "#SBATCH --nodes=1", + "#SBATCH --ntasks=1", + f"#SBATCH --array=0-{self.nodes - 1}", + ] + ) + else: + script_lines.extend( + [ + f"#SBATCH --output={self.output_dir}/madengine-{model_info['name']}_%j_%t.out", + f"#SBATCH --error={self.output_dir}/madengine-{model_info['name']}_%j_%t.err", + f"#SBATCH --partition={self.partition}", + f"#SBATCH --nodes={self.nodes}", + f"#SBATCH --ntasks={self.nodes}", + ] + ) if not self.skip_gpus_directive: script_lines.append(f"#SBATCH --gpus-per-node={self.gpus_per_node}") - script_lines += [ - f"#SBATCH --time={self.time_limit}", - ] + script_lines.extend( + [ + f"#SBATCH --time={self.time_limit}", + ] + ) # Honour user-configured exclusivity (defaults to True to match the standard SLURM template). if self.slurm_config.get("exclusive", True): script_lines.append("#SBATCH --exclusive") - + # Add reservation if specified if self.reservation: script_lines.append(f"#SBATCH --reservation={self.reservation}") - + # Add nodelist if specified (from model card or --additional-context) nodelist = self._normalize_nodelist(self.slurm_config.get("nodelist")) if nodelist: script_lines.append(f"#SBATCH --nodelist={nodelist}") - - script_lines.extend([ - "", - f"# slurm_multi launcher script for {model_info['name']}", - f"# Generated by madengine for slurm_multi", - "", - "set -e", - "", - "# Environment variables", - ]) - + + script_lines.extend( + [ + "", + f"# slurm_multi launcher script for {model_info['name']}", + f"# Generated by madengine for slurm_multi", + "", + "set -e", + "", + "# Environment variables", + ] + ) + for key, value in env_vars.items(): script_lines.append(f"export {key}={shlex.quote(str(value))}") - + script_lines.append("") - script_lines.extend([ - "echo '=========================================='", - "echo 'slurm_multi Launcher'", - "echo '=========================================='", - f"echo 'Model: {model_info['name']}'", - f"echo 'Script: {model_script_path}'", - "echo 'SLURM_JOB_ID:' $SLURM_JOB_ID", - "echo 'SLURM_NNODES:' $SLURM_NNODES", - "echo 'SLURM_NODELIST:' $SLURM_NODELIST", - "echo ''", - ]) - + if self.IS_SPUR: + # spur job-array rank + shared-filesystem rendezvous. Each array task + # is one node: SLURM_ARRAY_TASK_ID is the node rank. Pin SLURM_JOB_ID to + # the shared SLURM_ARRAY_JOB_ID so the launcher's rendezvous port and + # /run_logs/ dir match across nodes. rank 0 publishes its transport + # IP; peers read it as MASTER_ADDR (see also job.sh.j2 spur branch). + # Imported lazily: spur.SpurDeployment subclasses this module. + from .spur import DEFAULT_RENDEZVOUS_TIMEOUT, render_rendezvous_block + + rendezvous_dir = getattr( + self, + "rendezvous_dir", + str(self.output_dir.resolve() / "spur_rendezvous"), + ) + rendezvous_timeout = getattr( + self, "rendezvous_timeout", DEFAULT_RENDEZVOUS_TIMEOUT + ) + script_lines.extend( + [ + "# --- spur job-array rank + rendezvous ---", + 'export NODE_RANK="${SLURM_ARRAY_TASK_ID:-0}"', + 'export SLURM_PROCID="${SLURM_ARRAY_TASK_ID:-0}"', + f"export NNODES={self.nodes}", + f"export SLURM_NNODES={self.nodes}", + f"export WORLD_SIZE={self.nodes}", + 'export SLURM_JOB_ID="${SLURM_ARRAY_JOB_ID:-$SLURM_JOB_ID}"', + f'export SLURM_SUBMIT_DIR="${{SLURM_SUBMIT_DIR:-{manifest_dir}}}"', + ] + ) + script_lines.extend( + render_rendezvous_block(rendezvous_dir, rendezvous_timeout) + ) + script_lines.append("") + script_lines.extend( + [ + "echo '=========================================='", + "echo 'slurm_multi Launcher'", + "echo '=========================================='", + f"echo 'Model: {model_info['name']}'", + f"echo 'Script: {model_script_path}'", + "echo 'SLURM_JOB_ID:' $SLURM_JOB_ID", + "echo 'SLURM_NNODES:' $SLURM_NNODES", + "echo 'SLURM_NODELIST:' $SLURM_NODELIST", + "echo ''", + ] + ) + # Check if image needs parallel pull on all nodes # Pull if: image is from registry (contains / or .) and not a local ci-* build docker_image = env_vars.get("DOCKER_IMAGE_NAME", "") - is_registry_image = docker_image and not docker_image.startswith("ci-") and ("/" in docker_image or "." in docker_image) - - if is_registry_image: + is_registry_image = ( + docker_image + and not docker_image.startswith("ci-") + and ("/" in docker_image or "." in docker_image) + ) + + if is_registry_image and self.IS_SPUR: + # spur: each array task is its own node, so pull locally (no srun fan-out). + script_lines.extend( + [ + "", + "# Pull Docker image on this node (one array task per node)", + f"MAD_PULL_IMAGE={shlex.quote(docker_image)}", + "echo '=========================================='", + 'echo "[$(hostname)] Pulling $MAD_PULL_IMAGE..."', + "echo '=========================================='", + 'docker pull "$MAD_PULL_IMAGE"', + "PULL_EXIT=$?", + "if [ $PULL_EXIT -ne 0 ]; then", + ' echo "[$(hostname)] Docker pull failed for $MAD_PULL_IMAGE"', + " exit $PULL_EXIT", + "fi", + "echo ''", + ] + ) + elif is_registry_image: # Add parallel docker pull on all nodes # This ensures all nodes have the image before running - script_lines.extend([ - "", - "# Pull Docker image in parallel on all nodes", - "echo '=========================================='", - "echo 'Pulling Docker image on all nodes in parallel'", - "echo '=========================================='", - f"echo 'Image: {docker_image}'", - "echo ''", - "", - f"srun --nodes=$SLURM_NNODES --ntasks=$SLURM_NNODES bash -c \"", - f" echo \\\"[\\$(hostname)] Pulling {docker_image}...\\\"", - f" docker pull {docker_image}", - " PULL_RC=\\$?", - " if [ \\$PULL_RC -eq 0 ]; then", - " echo \\\"[\\$(hostname)] Pull SUCCESS\\\"", - " else", - " echo \\\"[\\$(hostname)] Pull FAILED with exit code \\$PULL_RC\\\"", - " fi", - " exit \\$PULL_RC", - "\"", - "PULL_EXIT=$?", - "", - "if [ $PULL_EXIT -ne 0 ]; then", - " echo 'Docker pull failed on one or more nodes'", - " exit $PULL_EXIT", - "fi", - "", - "echo ''", - "echo 'Docker image pulled on all nodes'", - "echo ''", - ]) - + script_lines.extend( + [ + "", + "# Pull Docker image in parallel on all nodes", + "echo '=========================================='", + "echo 'Pulling Docker image on all nodes in parallel'", + "echo '=========================================='", + f"echo 'Image: {docker_image}'", + "echo ''", + "", + f'srun --nodes=$SLURM_NNODES --ntasks=$SLURM_NNODES bash -c "', + f' echo \\"[\\$(hostname)] Pulling {docker_image}...\\"', + f" docker pull {docker_image}", + " PULL_RC=\\$?", + " if [ \\$PULL_RC -eq 0 ]; then", + ' echo \\"[\\$(hostname)] Pull SUCCESS\\"', + " else", + ' echo \\"[\\$(hostname)] Pull FAILED with exit code \\$PULL_RC\\"', + " fi", + " exit \\$PULL_RC", + '"', + "PULL_EXIT=$?", + "", + "if [ $PULL_EXIT -ne 0 ]; then", + " echo 'Docker pull failed on one or more nodes'", + " exit $PULL_EXIT", + "fi", + "", + "echo ''", + "echo 'Docker image pulled on all nodes'", + "echo ''", + ] + ) + # Create completion marker path for robust completion detection. # Namespace by SLURM_JOB_ID so concurrent / repeat runs of the same model # tag don't collide on each other's marker files. monitor() reconstructs # the same path using the deployment_id returned by sbatch. completion_marker_dir = self.output_dir.resolve() + # On spur every array task pins SLURM_JOB_ID to the shared SLURM_ARRAY_JOB_ID, + # so the job id alone is not unique per node - add the array rank. + marker_rank_suffix = "_rank${SLURM_ARRAY_TASK_ID:-0}" if self.IS_SPUR else "" completion_marker_template = ( completion_marker_dir - / f"madengine_{model_info['name']}_${{SLURM_JOB_ID:-local}}.complete" + / f"madengine_{model_info['name']}_${{SLURM_JOB_ID:-local}}{marker_rank_suffix}.complete" ) - + # Disable `set -e` around the model script bash invocation below so a # non-zero exit doesn't terminate the wrapper before SCRIPT_EXIT_CODE is # captured and the completion marker is written. monitor() relies on the # marker to distinguish 'failed' from 'still running'; without this, # a failed model run would look like a hang. - script_lines.extend([ - "", - "# Change to script directory", - f"cd {model_script_path.parent}", - "", - "# Run the model script directly on the host (with -e disabled so we", - "# can capture the exit code and write the completion marker even on failure).", - f"echo 'Executing: {_bash_invocation}'", - "set +e", - _bash_invocation, - "SCRIPT_EXIT_CODE=$?", - "set -e", - "", - "echo ''", - "echo 'Script completed.'", - "", - "# Write completion marker for madengine to detect (job-id namespaced)", - f"echo \"exit_code=$SCRIPT_EXIT_CODE\" > {completion_marker_template}", - f"echo \"timestamp=$(date -Iseconds)\" >> {completion_marker_template}", - f"echo \"Completion marker written: {completion_marker_template}\"", - "", - "exit $SCRIPT_EXIT_CODE", - ]) - + script_lines.extend( + [ + "", + "# Change to script directory", + f"cd {model_script_path.parent}", + "", + "# Run the model script directly on the host (with -e disabled so we", + "# can capture the exit code and write the completion marker even on failure).", + f"echo 'Executing: {_bash_invocation}'", + "set +e", + _bash_invocation, + "SCRIPT_EXIT_CODE=$?", + "set -e", + "", + "echo ''", + "echo 'Script completed.'", + "", + "# Write completion marker for madengine to detect (job-id namespaced)", + f'echo "exit_code=$SCRIPT_EXIT_CODE" > {completion_marker_template}', + f'echo "timestamp=$(date -Iseconds)" >> {completion_marker_template}', + f'echo "Completion marker written: {completion_marker_template}"', + ] + ) + if self.IS_SPUR: + # Per-rank marker consumed by SpurDeployment.monitor() (spur `sacct -j` + # is unreliable, so completion is detected from these files). + script_lines.extend( + [ + 'echo "$SCRIPT_EXIT_CODE" > "$_MAD_REND_DIR/done_rank${NODE_RANK}"', + ] + ) + script_lines.extend( + [ + "", + "exit $SCRIPT_EXIT_CODE", + ] + ) + # Store marker info for monitor() to reconstruct the path with deployment_id. self._completion_marker_dir = completion_marker_dir self._completion_marker_basename_template = ( @@ -641,32 +778,39 @@ def _prepare_slurm_multi_script(self, model_info: Dict, docker_image_name: str = self._completion_marker = ( completion_marker_dir / f"madengine_{model_info['name']}_local.complete" ) - + script_content = "\n".join(script_lines) - + # Save script self.script_path = self.output_dir / f"madengine_{model_info['name']}.sh" self.script_path.write_text(script_content) self.script_path.chmod(0o755) - - self.console.print(f"[green]✓ Generated slurm_multi script: {self.script_path}[/green]") + + self.console.print( + f"[green]✓ Generated slurm_multi script: {self.script_path}[/green]" + ) self.console.print(f" Model script: {model_script_path}") self.console.print(f" Environment: {len(env_vars)} variables") - + return True + def _prepare_template_context(self, model_info: Dict) -> Dict[str, Any]: """Prepare context for Jinja2 template rendering.""" # Use hierarchical GPU resolution: runtime > deployment > model > default additional_context = self.config.additional_context.copy() additional_context["slurm"] = self.slurm_config resolved_gpus_per_node = resolve_runtime_gpus(model_info, additional_context) - + # Extract launcher configuration - launcher_type = self.distributed_config.get("launcher", "torchrun") # Default to torchrun - + launcher_type = self.distributed_config.get( + "launcher", "torchrun" + ) # Default to torchrun + # Canonicalize aliases before validity check so e.g. sglang_disagg → sglang-disagg # passes through normalize_launcher instead of being mapped to "docker". - launcher_type = canonicalize_distributed_launcher(launcher_type) or launcher_type + launcher_type = ( + canonicalize_distributed_launcher(launcher_type) or launcher_type + ) # Normalize launcher based on deployment type and validity launcher_type = normalize_launcher(launcher_type, "slurm") # Persist the resolved launcher so downstream readers (reporting paths, @@ -675,9 +819,11 @@ def _prepare_template_context(self, model_info: Dict) -> Dict[str, Any]: self.distributed_config["launcher"] = launcher_type nnodes = self.distributed_config.get("nnodes", self.nodes) - nproc_per_node = self.distributed_config.get("nproc_per_node", resolved_gpus_per_node) + nproc_per_node = self.distributed_config.get( + "nproc_per_node", resolved_gpus_per_node + ) master_port = self.distributed_config.get("port", 29500) - + # Apply multi-node profiling logic if tools are configured tools = additional_context.get("tools", []) if nnodes > 1 and tools: @@ -686,28 +832,29 @@ def _prepare_template_context(self, model_info: Dict) -> Dict[str, Any]: class ConsoleLogger: def __init__(self, console): self.console = console + def info(self, msg): self.console.print(f"[cyan]{msg}[/cyan]") + def warning(self, msg): self.console.print(f"[yellow]{msg}[/yellow]") + def debug(self, msg): pass # Skip debug messages in console - + profiling_config = configure_multi_node_profiling( - nnodes=nnodes, - tools_config=tools, - logger=ConsoleLogger(self.console) + nnodes=nnodes, tools_config=tools, logger=ConsoleLogger(self.console) ) - + if profiling_config["enabled"]: tools = profiling_config["tools"] else: # rocprofv3 not available - skip profiling for multi-node tools = [] - + # Update tools in additional_context additional_context["tools"] = tools - + # Generate launcher-specific command launcher_command = self._generate_launcher_command( launcher_type=launcher_type, @@ -716,7 +863,7 @@ def debug(self, msg): master_port=master_port, model_name=model_info.get("name", "") or "", ) - + return { "model_name": model_info["name"], "manifest_file": os.path.abspath(self.config.manifest_file), @@ -748,9 +895,9 @@ def debug(self, msg): "live_output": self.config.additional_context.get("live_output", False), "tags": " ".join(model_info.get("tags", [])), "multiple_results": model_info.get("multiple_results"), - "credential_file": "credential.json" - if Path("credential.json").exists() - else None, + "credential_file": ( + "credential.json" if Path("credential.json").exists() else None + ), "data_file": "data.json" if Path("data.json").exists() else None, # Launcher configuration "launcher_type": launcher_type, @@ -759,6 +906,9 @@ def debug(self, msg): "nproc_per_node": nproc_per_node, # Profiling tools (processed for multi-node compatibility) "tools": tools, + # Scheduler flavor: "slurm" (default) or "spur" (job-array fan-out). + # SpurDeployment overrides this and adds "rendezvous_dir". + "scheduler": "slurm", } def _generate_launcher_command( @@ -771,15 +921,15 @@ def _generate_launcher_command( ) -> str: """ Generate launcher-specific command based on launcher type. - + Follows k8s pattern: different launchers have different command generation. - + Args: launcher_type: Type of launcher (torchrun, vllm, sglang, deepspeed, etc.) nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master communication port - + Returns: Launcher-specific environment setup and command string """ @@ -790,13 +940,17 @@ def _generate_launcher_command( elif launcher_type == "sglang": return self._generate_sglang_command(nnodes, nproc_per_node, master_port) elif launcher_type == "sglang-disagg": - return self._generate_sglang_disagg_command(nnodes, nproc_per_node, master_port) + return self._generate_sglang_disagg_command( + nnodes, nproc_per_node, master_port + ) elif launcher_type == "deepspeed": return self._generate_deepspeed_command(nnodes, nproc_per_node, master_port) elif launcher_type == "megatron": return self._generate_megatron_command(nnodes, nproc_per_node, master_port) elif launcher_type == "torchtitan": - return self._generate_torchtitan_command(nnodes, nproc_per_node, master_port) + return self._generate_torchtitan_command( + nnodes, nproc_per_node, master_port + ) elif launcher_type == "primus": return self._generate_primus_command( nnodes, nproc_per_node, master_port, model_name=model_name @@ -815,15 +969,15 @@ def _generate_torchrun_command( ) -> str: """ Generate torchrun launcher command for SLURM. - + For single-node (nnodes=1): Uses standalone mode For multi-node (nnodes>1): Uses distributed mode with SLURM environment - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: MAD_MULTI_NODE_RUNNER environment variable setup """ @@ -839,83 +993,83 @@ def _generate_vllm_command( ) -> str: """ Generate vLLM launcher environment variables. - + vLLM manages its own process spawning - no torchrun needed. Model script directly invokes vLLM with tensor/pipeline parallelism. - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: Environment variable setup for vLLM """ if nnodes == 1: - return f'''# vLLM single-node setup (Tensor Parallelism) + return f"""# vLLM single-node setup (Tensor Parallelism) export VLLM_TENSOR_PARALLEL_SIZE={nproc_per_node} export VLLM_PIPELINE_PARALLEL_SIZE=1 export VLLM_DISTRIBUTED_BACKEND="auto" -# vLLM handles its own process management - no MAD_MULTI_NODE_RUNNER needed''' +# vLLM handles its own process management - no MAD_MULTI_NODE_RUNNER needed""" else: # One vLLM serve per node (TP only on that node), no shared Ray = data parallelism - return f'''# vLLM multi-node setup (data parallel: one serve per node, TP only) + return f"""# vLLM multi-node setup (data parallel: one serve per node, TP only) export VLLM_TENSOR_PARALLEL_SIZE={nproc_per_node} export VLLM_PIPELINE_PARALLEL_SIZE=1 export VLLM_DISTRIBUTED_BACKEND="none" -# vLLM handles its own process management - no MAD_MULTI_NODE_RUNNER needed''' +# vLLM handles its own process management - no MAD_MULTI_NODE_RUNNER needed""" def _generate_sglang_command( self, nnodes: int, nproc_per_node: int, master_port: int ) -> str: """ Generate SGLang launcher environment variables. - + SGLang similar to vLLM - manages its own process spawning. - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: Environment variable setup for SGLang """ if nnodes == 1: - return f'''# SGLang single-node setup (Tensor Parallelism) + return f"""# SGLang single-node setup (Tensor Parallelism) export SGLANG_TENSOR_PARALLEL_SIZE={nproc_per_node} export SGLANG_PIPELINE_PARALLEL_SIZE=1 -# SGLang handles its own process management - no MAD_MULTI_NODE_RUNNER needed''' +# SGLang handles its own process management - no MAD_MULTI_NODE_RUNNER needed""" else: # One SGLang serve per node (TP only on that node), no cross-node coordination = data parallel - return f'''# SGLang multi-node setup (data parallel: one serve per node, TP only) + return f"""# SGLang multi-node setup (data parallel: one serve per node, TP only) export SGLANG_TENSOR_PARALLEL_SIZE={nproc_per_node} export SGLANG_PIPELINE_PARALLEL_SIZE=1 -# SGLang handles its own process management - no MAD_MULTI_NODE_RUNNER needed''' +# SGLang handles its own process management - no MAD_MULTI_NODE_RUNNER needed""" def _generate_sglang_disagg_command( self, nnodes: int, nproc_per_node: int, master_port: int ) -> str: """ Generate SGLang Disaggregated launcher environment for SLURM. - + SGLang Disaggregated Architecture: - Prefill nodes: xP - Decode nodes: yD - Proxy/router: a dedicated node (1 + xP + yD == nnodes) or co-located on the first prefill node (xP + yD == nnodes). The rank-to-role assignment is handled by the model run.sh, not by this launcher. - + Minimum cluster: 2 nodes (co-located proxy + 1 prefill + 1 decode) - + Args: nnodes: Total number of nodes (must be >= 2) nproc_per_node: GPUs per node (tensor parallel size) master_port: Master port for coordination - + Returns: Environment setup with node role assignment - + Raises: ValueError: If nnodes < 2 (minimum for disagg) """ @@ -924,12 +1078,14 @@ def _generate_sglang_disagg_command( f"SGLang Disaggregated requires minimum 2 nodes " f"(co-located proxy + 1 prefill + 1 decode), got {nnodes}" ) - + # Check if custom split is specified in additional_context - sglang_disagg_config = self.config.additional_context.get("distributed", {}).get("sglang_disagg", {}) + sglang_disagg_config = self.config.additional_context.get( + "distributed", {} + ).get("sglang_disagg", {}) prefill_nodes = sglang_disagg_config.get("prefill_nodes") decode_nodes = sglang_disagg_config.get("decode_nodes") - + if prefill_nodes is not None and decode_nodes is not None: # User specified custom split - validate if prefill_nodes < 1 or decode_nodes < 1: @@ -941,8 +1097,10 @@ def _generate_sglang_disagg_command( # co-located proxy/router on the first prefill node (xP + yD == nnodes), # mirroring the vllm-disagg layout. The proxy is launched by the model # script (not by this launcher), so both topologies are valid here. - if (prefill_nodes + decode_nodes != nnodes - and prefill_nodes + decode_nodes + 1 != nnodes): + if ( + prefill_nodes + decode_nodes != nnodes + and prefill_nodes + decode_nodes + 1 != nnodes + ): raise ValueError( f"Custom split validation failed: prefill_nodes ({prefill_nodes}) + " f"decode_nodes ({decode_nodes}) = {prefill_nodes + decode_nodes} must equal " @@ -962,8 +1120,8 @@ def _generate_sglang_disagg_command( # For N total nodes: 1 proxy + ~40% prefill + ~60% decode xP = max(1, (nnodes - 1) * 2 // 5) # ~40% of worker nodes yD = nnodes - 1 - xP # remaining nodes - - return f'''# SGLang Disaggregated multi-node setup + + return f"""# SGLang Disaggregated multi-node setup # ============================================ # Cluster Configuration: # Total Nodes: {nnodes} @@ -983,13 +1141,27 @@ def _generate_sglang_disagg_command( # Master coordination export MASTER_PORT={master_port} -# Build node IP list from SLURM -SLURM_NODE_IPS=$(scontrol show hostname ${{SLURM_JOB_NODELIST}} | while read node; do +# Build node IP list from SLURM (scheduler-portable). +# Stock SLURM: SLURM_JOB_NODELIST is compressed -> expand with `scontrol show +# hostnames`. Spur: it is already a comma-separated expanded list and +# `scontrol show hostname[s]` is unsupported. Prefer the mad_expand_nodelist +# helper defined in the job script header when available. +if command -v mad_expand_nodelist >/dev/null 2>&1; then + _MAD_NODES=$(mad_expand_nodelist "${{SLURM_JOB_NODELIST}}") +elif [[ "${{SLURM_JOB_NODELIST}}" != *"["* ]]; then + _MAD_NODES=$(echo "${{SLURM_JOB_NODELIST}}" | tr ',' '\\n') +else + _MAD_NODES=$(scontrol show hostnames "${{SLURM_JOB_NODELIST}}" 2>/dev/null || echo "${{SLURM_JOB_NODELIST}}" | tr ',' '\\n') +fi +SLURM_NODE_IPS=$(echo "$_MAD_NODES" | while read node; do + [ -z "$node" ] && continue getent hosts "$node" | awk '{{print $1}}' done | tr '\\n' ',' | sed 's/,$//') export SGLANG_NODE_IPS="$SLURM_NODE_IPS" -export SGLANG_NODE_RANK=${{SLURM_PROCID}} +# Node rank: stock SLURM sets SLURM_PROCID per srun task. Spur leaves it empty +# and instead provides NODE_RANK / PMI_RANK / SLURM_NODEID. Fall back in order. +export SGLANG_NODE_RANK=${{SLURM_PROCID:-${{NODE_RANK:-${{PMI_RANK:-${{SLURM_NODEID:-0}}}}}}}} echo "==========================================" echo "SGLang Disaggregated Cluster Info" @@ -1002,21 +1174,21 @@ def _generate_sglang_disagg_command( echo "==========================================" # No MAD_MULTI_NODE_RUNNER - SGLang disagg handles process management -# Model script should detect SGLANG_DISAGG_MODE and launch appropriately''' +# Model script should detect SGLANG_DISAGG_MODE and launch appropriately""" def _generate_deepspeed_command( self, nnodes: int, nproc_per_node: int, master_port: int ) -> str: """ Generate DeepSpeed launcher command. - + DeepSpeed has its own launcher similar to torchrun. - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: MAD_MULTI_NODE_RUNNER with deepspeed launcher """ @@ -1025,9 +1197,16 @@ def _generate_deepspeed_command( export MAD_MULTI_NODE_RUNNER="deepspeed --num_gpus={nproc_per_node}"''' else: return f'''# DeepSpeed multi-node setup -# Generate hostfile dynamically from SLURM +# Generate hostfile dynamically from SLURM (scheduler-portable nodelist expansion). +if command -v mad_expand_nodelist >/dev/null 2>&1; then + _MAD_DS_NODES=$(mad_expand_nodelist "$SLURM_JOB_NODELIST") +elif [[ "$SLURM_JOB_NODELIST" != *"["* ]]; then + _MAD_DS_NODES=$(echo "$SLURM_JOB_NODELIST" | tr ',' '\\n') +else + _MAD_DS_NODES=$(scontrol show hostnames "$SLURM_JOB_NODELIST" 2>/dev/null || echo "$SLURM_JOB_NODELIST" | tr ',' '\\n') +fi cat > /tmp/deepspeed_hostfile_${{SLURM_JOB_ID}}.txt << EOF -$(scontrol show hostnames $SLURM_JOB_NODELIST | awk -v slots={nproc_per_node} '{{print $1" slots="slots}}') +$(echo "$_MAD_DS_NODES" | awk -v slots={nproc_per_node} 'NF{{print $1" slots="slots}}') EOF export MAD_MULTI_NODE_RUNNER="deepspeed --hostfile=/tmp/deepspeed_hostfile_${{SLURM_JOB_ID}}.txt --master_addr=${{MASTER_ADDR}} --master_port={master_port}"''' @@ -1036,14 +1215,14 @@ def _generate_megatron_command( ) -> str: """ Generate Megatron-LM launcher command. - + Megatron-LM typically uses torchrun but with specific environment variables. - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: MAD_MULTI_NODE_RUNNER with megatron-specific setup """ @@ -1066,24 +1245,24 @@ def _generate_torchtitan_command( ) -> str: """ Generate TorchTitan launcher command for SLURM. - + TorchTitan is a PyTorch native platform for LLM pre-training that uses torchrun as its underlying launcher but requires additional configuration for multi-dimensional parallelism (FSDP2, Tensor Parallel, Pipeline Parallel). - + Key TorchTitan features: - Uses TOML configuration files for training setup - Supports FSDP2, Tensor Parallel, Pipeline Parallel, Context Parallel - Built on top of torchrun for distributed coordination - + For single-node (nnodes=1): Uses standalone torchrun mode For multi-node (nnodes>1): Uses distributed torchrun with SLURM environment - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: MAD_MULTI_NODE_RUNNER with torchtitan-specific setup """ @@ -1121,15 +1300,21 @@ def _generate_primus_command( We only export PRIMUS_CONFIG_PATH and optional PRIMUS_CLI_EXTRA. No MAD_MULTI_NODE_RUNNER. """ primus_cfg = merged_primus_config( - self.manifest if isinstance(getattr(self, "manifest", None), dict) else None, + ( + self.manifest + if isinstance(getattr(self, "manifest", None), dict) + else None + ), self.config.additional_context, ) config_path = primus_cfg.get("config_path", "exp_pretrain.yaml") cli_extra = primus_cfg.get("cli_extra", "") # Safe shell quoting for config_path and cli_extra config_path_quoted = config_path.replace('"', '\\"') - lines = [f'# Primus launcher (model script runs run_pretrain.sh)', - f'export PRIMUS_CONFIG_PATH="{config_path_quoted}"'] + lines = [ + f"# Primus launcher (model script runs run_pretrain.sh)", + f'export PRIMUS_CONFIG_PATH="{config_path_quoted}"', + ] if (cli_extra or "").strip(): cli_extra_quoted = cli_extra.replace('"', '\\"') lines.append(f'export PRIMUS_CLI_EXTRA="{cli_extra_quoted}"') @@ -1147,23 +1332,23 @@ def _generate_basic_env_command( ) -> str: """ Generate basic environment variables for unknown launchers. - + Provides standard distributed execution environment variables and lets the model script handle launcher invocation. - + Args: nnodes: Number of nodes nproc_per_node: GPUs per node master_port: Master port - + Returns: Basic environment variable setup """ - return f'''# Basic distributed environment (custom launcher) + return f"""# Basic distributed environment (custom launcher) export NNODES={nnodes} export NPROC_PER_NODE={nproc_per_node} export MASTER_PORT={master_port} -# Model script should handle launcher invocation''' +# Model script should handle launcher invocation""" def deploy(self) -> DeploymentResult: """Submit sbatch script to SLURM scheduler (locally).""" @@ -1185,13 +1370,30 @@ def deploy(self) -> DeploymentResult: # Single-node: we still run the check so bad nodes (e.g. Docker broken) get excluded; # we never gate submission for nodes==1 so behavior stays backward compatible. # Health-check srun invocations create SLURM jobs; we cancel them after preflight. - enable_preflight = self.slurm_config.get("enable_node_check", True) + # Not applicable to spur: the check runs `srun -w `, which spur + # executes on the head node instead of the requested one, and it pins the + # job to a multi-node `#SBATCH --nodelist`. Under the spur job array every + # task requests `--nodes=1`, so a multi-node nodelist would make each task + # demand every node (and, with --exclusive, deadlock the rendezvous). + enable_preflight = ( + self.slurm_config.get("enable_node_check", True) and not self.IS_SPUR + ) auto_cleanup = self.slurm_config.get("auto_cleanup_nodes", False) - allow_submit_without_clean = self.slurm_config.get("allow_submit_without_clean_nodes", False) + allow_submit_without_clean = self.slurm_config.get( + "allow_submit_without_clean_nodes", False + ) clean_nodes: List[str] = [] health_check_job_name: Optional[str] = None + if self.IS_SPUR and self.slurm_config.get("enable_node_check"): + self.console.print( + "[dim]Node health preflight is not supported on spur; skipping.[/dim]" + ) - if enable_preflight and self.nodes >= 1 and not self.slurm_config.get("nodelist"): + if ( + enable_preflight + and self.nodes >= 1 + and not self.slurm_config.get("nodelist") + ): try: selector = SlurmNodeSelector( console=self.console, @@ -1205,10 +1407,14 @@ def deploy(self) -> DeploymentResult: exclude=self.slurm_config.get("exclude"), constraint=self.slurm_config.get("constraint"), ) - health_check_job_name = getattr(selector, "_health_check_job_name", None) + health_check_job_name = getattr( + selector, "_health_check_job_name", None + ) # Update exclude list if we found dirty/unreachable/unknown nodes - if updated_exclude and updated_exclude != self.slurm_config.get("exclude", ""): + if updated_exclude and updated_exclude != self.slurm_config.get( + "exclude", "" + ): self.console.print( f"[dim]Updated exclude list for sbatch: {updated_exclude}[/dim]\n" ) @@ -1221,7 +1427,9 @@ def deploy(self) -> DeploymentResult: and not allow_submit_without_clean and len(clean_nodes) < self.nodes ): - SlurmNodeSelector.cancel_health_check_jobs(health_check_job_name, self.console) + SlurmNodeSelector.cancel_health_check_jobs( + health_check_job_name, self.console + ) return DeploymentResult( status=DeploymentStatus.FAILED, deployment_id="", @@ -1238,23 +1446,49 @@ def deploy(self) -> DeploymentResult: self.console.print(f"[dim]Using nodelist: {nodelist_str}[/dim]\n") self.prepare() except Exception as e: - self.console.print( - f"[yellow]⚠ Node health check failed: {e}[/yellow]" - ) + self.console.print(f"[yellow]⚠ Node health check failed: {e}[/yellow]") self.console.print("[dim]Continuing with job submission[/dim]\n") finally: # Always cancel health-check jobs so they do not stay in the queue - SlurmNodeSelector.cancel_health_check_jobs(health_check_job_name, self.console) + SlurmNodeSelector.cancel_health_check_jobs( + health_check_job_name, self.console + ) # ==================== END PREFLIGHT ==================== try: - # Submit job to SLURM (runs locally on login node) - result = subprocess.run( - ["sbatch", str(self.script_path)], - capture_output=True, - text=True, - timeout=30, + # Submit job to SLURM (runs locally on login node). + # Spur compatibility: spur's control plane is Raft-based and can + # transiently reject submissions with "not the Raft leader" / "no + # leader elected yet" during leader election. Retry a few times. + _SUBMIT_RETRIES = 6 + _SUBMIT_RETRY_DELAY = 5 # seconds + _TRANSIENT_MARKERS = ( + "not the raft leader", + "no leader elected", + "service is currently unavailable", + "leader elected yet", ) + result = None + for _attempt in range(1, _SUBMIT_RETRIES + 1): + result = subprocess.run( + ["sbatch", str(self.script_path)], + capture_output=True, + text=True, + timeout=30, + ) + if result.returncode == 0: + break + _err = (result.stderr or "").lower() + _is_transient = any(m in _err for m in _TRANSIENT_MARKERS) + if _is_transient and _attempt < _SUBMIT_RETRIES: + self.console.print( + f"[dim yellow]sbatch transient scheduler error " + f"(attempt {_attempt}/{_SUBMIT_RETRIES}), retrying in " + f"{_SUBMIT_RETRY_DELAY}s...[/dim yellow]" + ) + time.sleep(_SUBMIT_RETRY_DELAY) + continue + break if result.returncode == 0: # Parse job ID: "Submitted batch job 12345" @@ -1293,7 +1527,7 @@ def deploy(self) -> DeploymentResult: def _run_inside_existing_allocation(self) -> DeploymentResult: """ Run script directly inside existing salloc allocation using bash. - + The script will use the nodes already allocated to the current job. SLURM environment variables (SLURM_NODELIST, etc.) are inherited. """ @@ -1305,16 +1539,18 @@ def _run_inside_existing_allocation(self) -> DeploymentResult: deployment_id=self.existing_job_id, message=error_msg, ) - + self.console.print( f"\n[bold cyan]Running inside existing SLURM allocation[/bold cyan]" ) self.console.print(f" Job ID: {self.existing_job_id}") - self.console.print(f" Using {self.nodes} of {self.allocation_nodes} allocated nodes") + self.console.print( + f" Using {self.nodes} of {self.allocation_nodes} allocated nodes" + ) self.console.print(f" GPUs per node: {self.gpus_per_node}") self.console.print(f" Script: {self.script_path}") self.console.print(f"\n[dim]Executing: bash {self.script_path}[/dim]\n") - + try: # Run script directly with bash (synchronous, blocks until done) # Don't capture output - let it stream directly to console @@ -1322,7 +1558,7 @@ def _run_inside_existing_allocation(self) -> DeploymentResult: ["bash", str(self.script_path)], timeout=subprocess_timeout(self.config.timeout), ) - + if result.returncode == 0: self.console.print( f"\n[green]✓ Script completed successfully in allocation {self.existing_job_id}[/green]" @@ -1345,7 +1581,7 @@ def _run_inside_existing_allocation(self) -> DeploymentResult: logs_path=str(self.output_dir), skip_monitoring=True, # Already ran synchronously ) - + except subprocess.TimeoutExpired: self.console.print( f"\n[red]✗ Script timed out after {self.config.timeout}s[/red]" @@ -1379,7 +1615,7 @@ def monitor(self, deployment_id: str) -> DeploymentResult: return self._check_job_completion(deployment_id) status = result.stdout.strip().upper() - + # Check if live output is enabled live_output = self.config.additional_context.get("live_output", False) @@ -1418,8 +1654,11 @@ def monitor(self, deployment_id: str) -> DeploymentResult: ) except Exception as e: - self.console.print(f"[red]Monitor exception for job {deployment_id}: {e}[/red]") + self.console.print( + f"[red]Monitor exception for job {deployment_id}: {e}[/red]" + ) import traceback + self.console.print(f"[dim red]{traceback.format_exc()}[/dim red]") return DeploymentResult( status=DeploymentStatus.FAILED, @@ -1430,80 +1669,88 @@ def monitor(self, deployment_id: str) -> DeploymentResult: def _stream_job_output(self, job_id: str, final: bool = False): """Stream output from SLURM job output file.""" # Track last position read from output file - if not hasattr(self, '_output_positions'): + if not hasattr(self, "_output_positions"): self._output_positions = {} - + # Find output file output_dir = str(self.output_dir) output_pattern = f"{output_dir}/madengine-*_{job_id}_*.out" - + try: import glob + output_files = glob.glob(output_pattern) - + if not output_files: return # Output file not created yet - + output_file = output_files[0] # Use first match - + # Read new content from file try: - with open(output_file, 'r') as f: + with open(output_file, "r") as f: # Seek to last position last_pos = self._output_positions.get(job_id, 0) f.seek(last_pos) - + # Read new lines new_content = f.read() - + if new_content: # Print new output with prefix for line in new_content.splitlines(): if line.strip(): # Skip empty lines self.console.print(f"[dim cyan]│[/dim cyan] {line}") - + # Update position self._output_positions[job_id] = f.tell() - + except FileNotFoundError: pass # File not ready yet - + except Exception as e: # Silently ignore streaming errors to not disrupt monitoring if final: - self.console.print(f"[dim yellow]Note: Could not stream output: {e}[/dim yellow]") + self.console.print( + f"[dim yellow]Note: Could not stream output: {e}[/dim yellow]" + ) def _show_log_summary(self, job_id: str, success: bool = True): """Show a summary with pointers to log files instead of streaming verbose output.""" output_dir = str(self.output_dir) - + try: import glob + # Find output and error files for this job output_files = glob.glob(f"{output_dir}/madengine-*_{job_id}_*.out") error_files = glob.glob(f"{output_dir}/madengine-*_{job_id}_*.err") - + if output_files or error_files: status_symbol = "✓" if success else "✗" status_color = "green" if success else "red" - - self.console.print(f"[{status_color}]{status_symbol}[/{status_color}] SLURM job {job_id} logs saved to:") - + + self.console.print( + f"[{status_color}]{status_symbol}[/{status_color}] SLURM job {job_id} logs saved to:" + ) + for out_file in output_files: self.console.print(f" [cyan]→[/cyan] Output: {out_file}") - + for err_file in error_files: # Check if error file has content if os.path.exists(err_file) and os.path.getsize(err_file) > 0: self.console.print(f" [yellow]→[/yellow] Errors: {err_file}") - + if not success and error_files: # Show last few lines of error file for failed jobs for err_file in error_files: if os.path.exists(err_file) and os.path.getsize(err_file) > 0: - self.console.print(f"\n[yellow]Last 10 lines of error log:[/yellow]") + self.console.print( + f"\n[yellow]Last 10 lines of error log:[/yellow]" + ) try: - with open(err_file, 'r') as f: + with open(err_file, "r") as f: lines = f.readlines() for line in lines[-10:]: if line.strip(): @@ -1512,10 +1759,14 @@ def _show_log_summary(self, job_id: str, success: bool = True): pass break # Only show first error file else: - self.console.print(f"[dim yellow]Note: Log files for job {job_id} not found in {output_dir}[/dim yellow]") - + self.console.print( + f"[dim yellow]Note: Log files for job {job_id} not found in {output_dir}[/dim yellow]" + ) + except Exception as e: - self.console.print(f"[dim yellow]Note: Could not locate log files: {e}[/dim yellow]") + self.console.print( + f"[dim yellow]Note: Could not locate log files: {e}[/dim yellow]" + ) def _check_job_completion(self, job_id: str) -> DeploymentResult: """Check completed job status using sacct (locally). @@ -1549,11 +1800,13 @@ def _check_job_completion(self, job_id: str) -> DeploymentResult: if result.returncode == 0: status = result.stdout.strip().upper() - self.console.print(f"[dim]SLURM job {job_id} final status: {status}[/dim]") - + self.console.print( + f"[dim]SLURM job {job_id} final status: {status}[/dim]" + ) + # Check if live output is enabled live_output = self.config.additional_context.get("live_output", False) - + if "COMPLETED" in status: # Show final output or summary based on live_output flag if live_output: @@ -1619,9 +1872,13 @@ def _build_perf_entry_from_aggregated( run_details = { "model": model_info.get("name", aggregated_record.get("model", "")), - "n_gpus": str(aggregated_record.get("n_gpus", self.nodes * self.gpus_per_node)), + "n_gpus": str( + aggregated_record.get("n_gpus", self.nodes * self.gpus_per_node) + ), "nnodes": str(aggregated_record.get("nnodes", self.nodes)), - "gpus_per_node": str(aggregated_record.get("gpus_per_node", self.gpus_per_node)), + "gpus_per_node": str( + aggregated_record.get("gpus_per_node", self.gpus_per_node) + ), "training_precision": model_info.get("training_precision", ""), "pipeline": get_pipeline(), "args": model_info.get("args", ""), @@ -1646,7 +1903,9 @@ def _build_perf_entry_from_aggregated( "data_size": "", "data_download_duration": "", "build_number": get_build_number(), - "additional_docker_run_options": model_info.get("additional_docker_run_options", ""), + "additional_docker_run_options": model_info.get( + "additional_docker_run_options", "" + ), } flatten_tags(run_details) @@ -1701,12 +1960,16 @@ def _build_common_info_dict( "data_size": "", "data_download_duration": "", "build_number": get_build_number(), - "additional_docker_run_options": model_info.get("additional_docker_run_options", ""), + "additional_docker_run_options": model_info.get( + "additional_docker_run_options", "" + ), } flatten_tags(result) return result - def _select_best_multiple_results_csv(self, candidates: List[Path]) -> Optional[Path]: + def _select_best_multiple_results_csv( + self, candidates: List[Path] + ) -> Optional[Path]: """Pick the CSV with the most non-empty performance entries. In multi-node SLURM runs every node copies its local multi-results CSV @@ -1724,6 +1987,7 @@ def _select_best_multiple_results_csv(self, candidates: List[Path]) -> Optional[ if len(candidates) == 1: return candidates[0] import csv as _csv + best_candidate: Optional[Path] = None best_score = -1 best_rows = -1 @@ -1740,7 +2004,10 @@ def _select_best_multiple_results_csv(self, candidates: List[Path]) -> Optional[ for row in reader: total_rows += 1 if has_perf_column: - normalized_row = {(k.strip() if isinstance(k, str) else k): v for k, v in row.items()} + normalized_row = { + (k.strip() if isinstance(k, str) else k): v + for k, v in row.items() + } value = (normalized_row.get("performance") or "").strip() if value: non_empty_perf += 1 @@ -1759,7 +2026,6 @@ def _select_best_multiple_results_csv(self, candidates: List[Path]) -> Optional[ ) return best_candidate - def collect_results(self, deployment_id: str) -> Dict[str, Any]: """Collect performance results from SLURM output files. @@ -1793,7 +2059,9 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: built_models_dict = self.manifest.get("built_models") or {} model_info_for_path = built_models_dict.get(model_key, {}) if model_key else {} model_name_for_path = model_info_for_path.get("name", model_key or "unknown") - model_name = model_key or "unknown" # image key for build_info / model_info_for_entry lookups + model_name = ( + model_key or "unknown" + ) # image key for build_info / model_info_for_entry lookups # slurm_multi early dispatch: model script emits its own perf.csv directly, # so collect via _collect_slurm_multi_results instead of the template-based path. @@ -1805,7 +2073,6 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: deployment_id, results, session_start_row ) - build_info = {} built_images = self.manifest.get("built_images") or {} if built_images: @@ -1820,7 +2087,9 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: # Gather log content per node: from job_dir/node_N/ (new) or flat output_dir .out files per_node_log_contents: List[tuple] = [] - flat_out_files = sorted(self.output_dir.glob(f"madengine-*_{deployment_id}_*.out")) + flat_out_files = sorted( + self.output_dir.glob(f"madengine-*_{deployment_id}_*.out") + ) # Multi-node: only use explicit node logs (_node_N.out) to avoid also picking up # SBATCH %t output (madengine-*__0.out, _1.out), which would duplicate metrics. if self.nodes > 1: @@ -1851,7 +2120,9 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: # Multi-node: keep only log entries for actual node indices [0, nodes-1] if self.nodes > 1: - per_node_log_contents = [(n, c) for n, c in per_node_log_contents if n < self.nodes] + per_node_log_contents = [ + (n, c) for n, c in per_node_log_contents if n < self.nodes + ] # Copy flat logs into job_dir/node_/ for consistency if not already there. # Only create dirs for indices in [0, nodes-1] so we never create extra node_2, etc. @@ -1895,9 +2166,11 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: ) run_details_dict: Optional[Dict[str, Any]] = None - model_info_for_entry = (self.manifest.get("built_models") or {}).get( - model_key, {} - ) if model_key else {} + model_info_for_entry = ( + (self.manifest.get("built_models") or {}).get(model_key, {}) + if model_key + else {} + ) # Multiple results path: resolve CSV from job_dir/node_*, then cwd/run_directory mult_res = model_info_for_entry.get("multiple_results") @@ -1953,22 +2226,29 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: ) results["perf_files"] = [str(Path("perf.csv").resolve())] import csv as _csv + try: - with open(resolved_csv, "r", encoding="utf-8", errors="ignore") as f: + with open( + resolved_csv, "r", encoding="utf-8", errors="ignore" + ) as f: reader = _csv.DictReader(f) for row in reader: row = {k.strip(): v for k, v in row.items() if k} if row.get("performance") and row.get("metric"): - results["successful_runs"].append({ - "model": model_info_for_entry.get("name", "") + "_" + row.get("model", ""), - "status": "SUCCESS", - "performance": str(row.get("performance", "")), - "metric": row.get("metric", ""), - "duration": row.get("test_duration", ""), - "gpu_arch": gpu_arch, - "deployment": "slurm", - "machine": deployment_id, - }) + results["successful_runs"].append( + { + "model": model_info_for_entry.get("name", "") + + "_" + + row.get("model", ""), + "status": "SUCCESS", + "performance": str(row.get("performance", "")), + "metric": row.get("metric", ""), + "duration": row.get("test_duration", ""), + "gpu_arch": gpu_arch, + "deployment": "slurm", + "machine": deployment_id, + } + ) except Exception: pass self.console.print( @@ -2059,9 +2339,13 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: perf_csv_path = "perf.csv" self._ensure_perf_csv_exists() if run_details_dict.get("status") == "SUCCESS": - update_perf_csv(perf_csv=perf_csv_path, single_result=str(perf_entry_path)) + update_perf_csv( + perf_csv=perf_csv_path, single_result=str(perf_entry_path) + ) else: - update_perf_csv(perf_csv=perf_csv_path, exception_result=str(perf_entry_path)) + update_perf_csv( + perf_csv=perf_csv_path, exception_result=str(perf_entry_path) + ) try: scripts_path = model_info_for_entry.get("scripts", "") scripts_base_dir = scripts_base_dir_from(scripts_path) @@ -2083,7 +2367,9 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: num_entries=num_entries, ) except Exception as e: - self.console.print(f"[yellow]⚠ Could not update perf_super: {e}[/yellow]") + self.console.print( + f"[yellow]⚠ Could not update perf_super: {e}[/yellow]" + ) results["perf_files"] = [str(Path(perf_csv_path).resolve())] run_data = { "model": run_details_dict.get("model", ""), @@ -2127,17 +2413,22 @@ def collect_results(self, deployment_id: str) -> Dict[str, Any]: return results def _collect_slurm_multi_results( - self, deployment_id: str, results: Dict[str, Any], session_start_row: Optional[int] + self, + deployment_id: str, + results: Dict[str, Any], + session_start_row: Optional[int], ) -> Dict[str, Any]: """ Collect results for slurm_multi launchers. - + slurm_multi model scripts generate their own perf.csv via their benchmark scripts (e.g. generate_perf_csv.py). We collect SLURM logs for diagnostics and read the model-generated perf.csv for metrics. """ # Collect SLURM output logs for diagnostics - flat_out_files = sorted(self.output_dir.glob(f"madengine-*_{deployment_id}_*.out")) + flat_out_files = sorted( + self.output_dir.glob(f"madengine-*_{deployment_id}_*.out") + ) results["logs"] = [str(f) for f in flat_out_files] # Look for model-generated perf.csv. Inner scripts in MAD-private write @@ -2158,16 +2449,23 @@ def _collect_slurm_multi_results( user = os.environ.get("USER", "") shared_candidates = [] if user: - shared_candidates.extend([ - Path(f"/shared_inference/{user}/{deployment_id}/perf.csv"), - Path(f"/shared_inference/{user}/model_blog_logs/{deployment_id}/perf.csv"), - ]) + shared_candidates.extend( + [ + Path(f"/shared_inference/{user}/{deployment_id}/perf.csv"), + Path( + f"/shared_inference/{user}/model_blog_logs/{deployment_id}/perf.csv" + ), + ] + ) workspace_perf_dir = Path("slurm_output/perf_csv") - workspace_candidates = list(workspace_perf_dir.glob(f"*{deployment_id}*.csv")) + workspace_candidates = list( + workspace_perf_dir.glob(f"*{deployment_id}*.csv") + ) workspace_perf = Path("perf.csv") # Retry briefly for NFS propagation after SLURM job completion import time + for _attempt in range(6): for cand in shared_candidates: if cand.exists() and cand.stat().st_size > 0: @@ -2183,7 +2481,9 @@ def _collect_slurm_multi_results( break time.sleep(5) # Re-glob in case the file appeared during the wait. - workspace_candidates = list(workspace_perf_dir.glob(f"*{deployment_id}*.csv")) + workspace_candidates = list( + workspace_perf_dir.glob(f"*{deployment_id}*.csv") + ) if perf_csv_path: results["perf_files"] = [str(perf_csv_path)] @@ -2195,11 +2495,14 @@ def _collect_slurm_multi_results( # update_perf_csv(); slurm_multi flows did not, so this mirrors # that convention without modifying the original per-job file. import shutil + cwd_perf = Path("perf.csv") try: if cwd_perf.exists(): with open(perf_csv_path, "r") as src, open(cwd_perf, "a") as dst: - next(src, None) # skip per-job header so cwd CSV stays single-headed + next( + src, None + ) # skip per-job header so cwd CSV stays single-headed for line in src: dst.write(line) else: @@ -2212,7 +2515,9 @@ def _collect_slurm_multi_results( f"[yellow]⚠ Could not aggregate per-job perf into cwd perf.csv: {e}[/yellow]" ) else: - self.console.print("[yellow]No perf.csv found from slurm_multi model script[/yellow]") + self.console.print( + "[yellow]No perf.csv found from slurm_multi model script[/yellow]" + ) self.console.print( f"[green]Collected slurm_multi results: {len(results['perf_files'])} perf files, " @@ -2261,13 +2566,10 @@ def _collect_results_parse_perf_csv( def cleanup(self, deployment_id: str) -> bool: """Cancel SLURM job if still running (locally).""" try: - subprocess.run( - ["scancel", deployment_id], capture_output=True, timeout=10 - ) + subprocess.run(["scancel", deployment_id], capture_output=True, timeout=10) self.console.print(f"[yellow]Cancelled SLURM job: {deployment_id}[/yellow]") return True except Exception as e: self.console.print(f"[yellow]⚠ Cleanup warning: {e}[/yellow]") return False - diff --git a/src/madengine/deployment/spur.py b/src/madengine/deployment/spur.py new file mode 100644 index 00000000..7ca6a5ca --- /dev/null +++ b/src/madengine/deployment/spur.py @@ -0,0 +1,338 @@ +#!/usr/bin/env python3 +""" +Spur (Crusoe) deployment backend. + +Spur is an "AI-native" scheduler that exposes SLURM-compatible CLI shims +(sbatch/srun/squeue/sacct/scontrol/...) but differs from stock SLURM in ways +that break the standard multi-node flow: + + * `srun` cannot fan out tasks across nodes: any `srun [-N -n] [--mpi ...]` + invocation runs the command once on the head node, and SLURM_PROCID is + empty inside srun. The stock madengine template relies on + `srun bash task_script` launching one task per node with a unique + SLURM_PROCID, so only rank 0 would ever start. + * `scontrol show hostname[s]` is unsupported (SLURM_NODELIST is already an + expanded comma list). + * The control plane is Raft-based / eventually consistent: sbatch can + transiently fail ("not the Raft leader") and squeue/sacct states flap. + +Strategy: reuse the SLURM template and orchestration, but drive multi-node +execution with a job ARRAY of single-node tasks (one array task per node). +`SLURM_ARRAY_TASK_ID` is the node rank; the tasks self-form the cluster via the +model launcher's TCP rendezvous (rank 0 publishes its transport IP to a shared +filesystem, peers read it as MASTER_ADDR). The spur-specific branches live in +`templates/slurm/job.sh.j2` under `{% if scheduler == 'spur' %}` and are enabled +purely by the template context produced here. + +Copyright (c) Advanced Micro Devices, Inc. All rights reserved. +""" + +import os +import subprocess +from pathlib import Path +from typing import Any, Dict, List + +from .base import DeploymentConfig, DeploymentResult, DeploymentStatus +from .slurm import SlurmDeployment + +# Seconds a non-zero rank waits for rank 0 to publish MASTER_ADDR. A job array +# carries no gang-scheduling guarantee, so tasks can start minutes apart (and +# with --exclusive they may even start serially); override per site with +# slurm.rendezvous_timeout. +DEFAULT_RENDEZVOUS_TIMEOUT = 900 + + +def render_rendezvous_block(rendezvous_dir: str, timeout: int) -> List[str]: + """Bash lines resolving MASTER_ADDR via a shared-filesystem rendezvous. + + Rank 0 publishes its transport IP to ``//master_addr``; the other ranks poll for it. Requires ``SLURM_PROCID`` + to already hold the array task id; exports ``_MAD_REND_DIR`` (reused for + the per-rank ``done_rank`` markers) and ``MASTER_ADDR``. + + A peer that times out fails fast rather than continuing with an empty + MASTER_ADDR (which fails obscurely inside the launcher): it prints a + diagnostic, writes a non-zero ``done_rank`` marker so monitor() reports the + failure immediately, and exits non-zero. + + Args: + rendezvous_dir: Shared-filesystem root, visible from every node. + timeout: Seconds a non-zero rank waits for rank 0. + + Returns: + The bash lines, one per list element. + """ + return [ + f'_MAD_REND_DIR="{rendezvous_dir}/${{SLURM_ARRAY_JOB_ID:-${{SLURM_JOB_ID}}}}"', + 'mkdir -p "$_MAD_REND_DIR" 2>/dev/null || true', + '_MAD_IFACE="${NCCL_SOCKET_IFNAME:-ens3}"; _MAD_IFACE="${_MAD_IFACE%%,*}"', + '_MAD_MY_IP="$(ip -4 -o addr show "$_MAD_IFACE" 2>/dev/null | awk \'{print $4}\' | cut -d/ -f1 | head -n1)"', + '[ -z "$_MAD_MY_IP" ] && _MAD_MY_IP="$(hostname -I | awk \'{print $1}\')"', + 'if [ "${SLURM_PROCID}" = "0" ]; then', + ' echo "$_MAD_MY_IP" > "$_MAD_REND_DIR/master_addr"', + ' export MASTER_ADDR="$_MAD_MY_IP"', + "else", + f" _MAD_REND_TIMEOUT={int(timeout)}", + ' for _i in $(seq 1 "$_MAD_REND_TIMEOUT"); do [ -s "$_MAD_REND_DIR/master_addr" ] && break; sleep 1; done', + ' export MASTER_ADDR="$(cat "$_MAD_REND_DIR/master_addr" 2>/dev/null || true)"', + ' if [ -z "$MASTER_ADDR" ]; then', + ' echo "[spur-rendezvous] ERROR: rank ${SLURM_PROCID} on $(hostname) timed out after ${_MAD_REND_TIMEOUT}s waiting for $_MAD_REND_DIR/master_addr" >&2', + ' echo "[spur-rendezvous] Rank 0 never started (array tasks are not gang-scheduled), died early, or the rendezvous dir is not on a shared filesystem." >&2', + ' echo "[spur-rendezvous] Raise slurm.rendezvous_timeout above ${_MAD_REND_TIMEOUT}s if the queue wait is simply longer than that." >&2', + ' echo "1" > "$_MAD_REND_DIR/done_rank${SLURM_PROCID}"', + " exit 1", + " fi", + "fi", + 'echo "[spur-rendezvous] rank=${SLURM_PROCID} node=$(hostname) my_ip=$_MAD_MY_IP MASTER_ADDR=${MASTER_ADDR}"', + ] + + +class SpurDeployment(SlurmDeployment): + """SLURM-compatible deployment for the spur scheduler (job-array fan-out).""" + + DEPLOYMENT_TYPE = "spur" + # spur ships slurm-compatible shims. scontrol exists but is only partially + # implemented; the spur flow does not depend on it, so we don't require it. + REQUIRED_TOOLS = ["sbatch", "squeue", "sacct"] + # Drives the spur-specific branches in the inherited SLURM code paths + # (template rendering and the slurm_multi launcher): job-array fan-out + # instead of srun. + IS_SPUR = True + + def __init__(self, config: DeploymentConfig): + super().__init__(config) + # Rendezvous root MUST be on a shared (NFS) filesystem visible to every + # node: rank 0 writes MASTER_ADDR here and peers read it. output_dir is + # under the (shared) submission/run directory. + self.rendezvous_dir = str(self.output_dir.resolve() / "spur_rendezvous") + self.rendezvous_timeout = int( + self.slurm_config.get("rendezvous_timeout", DEFAULT_RENDEZVOUS_TIMEOUT) + ) + + def _prepare_template_context(self, model_info: Dict) -> Dict[str, Any]: + context = super()._prepare_template_context(model_info) + context["scheduler"] = "spur" + context["rendezvous_dir"] = self.rendezvous_dir + # Rendered as bash (not escaped) into the spur branch of job.sh.j2; the + # slurm_multi launcher script emits the same block. + context["spur_rendezvous_block"] = "\n".join( + render_rendezvous_block(self.rendezvous_dir, self.rendezvous_timeout) + ) + return context + + def _model_job_name(self) -> str: + """The #SBATCH --job-name used by the template (madengine-).""" + try: + models = self.manifest.get("built_models") or {} + first = next(iter(models.values()), {}) + name = first.get("name") or next(iter(models), "") + return f"madengine-{name}" + except Exception: + return "madengine-" + + def _live_task_count(self, deployment_id: str, job_name: str) -> int: + """Count my not-yet-finished array tasks in the queue (best-effort). + + Used only as a liveness guard so monitor() does not wait forever if a + task dies before writing its completion marker. squeue is eventually + consistent on spur, so a transient 0 is tolerated by the caller. + + Args: + deployment_id: Array job id returned by sbatch. + job_name: #SBATCH --job-name, used only if no row carries our id. + + Returns: + Number of live tasks, or -1 if squeue could not be queried. + """ + try: + # NOTE: spur's squeue ignores custom -o delimiters (e.g. "%j|%T" + # renders as " ", space-separated), so parse by + # whitespace. Job names produced by the template contain no spaces. + cmd = ["squeue", "-h", "-o", "%i %j %T"] + user = os.environ.get("USER", "") + if user: + # Omit -u entirely when USER is unset: `squeue -u ""` is an error. + cmd[1:1] = ["-u", user] + result = subprocess.run( + cmd, + capture_output=True, + text=True, + timeout=10, + ) + if result.returncode != 0: + return -1 # unknown + live_states = { + "PENDING", + "RUNNING", + "CONFIGURING", + "COMPLETING", + "RESIZING", + "SUSPENDED", + } + by_id = 0 + by_name = 0 + for line in result.stdout.splitlines(): + parts = line.split() + if len(parts) < 3: + continue + task_id, name, state = parts[0], parts[1], parts[-1] + if state.upper() not in live_states: + continue + # Array tasks are listed as "_" (or as a + # pending range, "_[1-3]"). Matching the id keeps a + # concurrent run of the same model from inflating the count. + if task_id == deployment_id or task_id.startswith(f"{deployment_id}_"): + by_id += 1 + elif name == job_name: + by_name += 1 + # Fall back to name matching only if squeue reported no row for our + # job id at all (i.e. it does not label array tasks the way we expect). + return by_id if by_id else by_name + except Exception: + return -1 # unknown + + # Number of consecutive polls, AFTER the tasks were first seen alive, with no + # completion markers AND no live tasks before we conclude the array died + # without reporting. ~poll interval (30s) times this many => grace window. + # The "seen alive first" gate is essential on spur: for the first ~1-2 min + # after sbatch, squeue does not yet list the array tasks (registration lag / + # eventual consistency), so a fresh, healthy run reports 0 live tasks. + _SPUR_DEAD_POLLS = 4 + + # Number of consecutive polls where squeue could not be queried at all before + # giving up. Completion markers still win if they appear, so this only bounds + # the case where the control plane stays unreachable and the markers never + # arrive; ~poll interval (30s) times this many => ~10 minutes. + _SPUR_UNKNOWN_POLLS = 20 + + # Number of consecutive polls with an empty queue BEFORE the tasks were ever + # seen alive. This bounds the startup grace window above: an array that fails + # before squeue ever lists it (bad partition, node failure, scheduler reject) + # writes no marker and never shows up, and the caller polls monitor() without + # a timeout. ~poll interval (30s) times this many => ~10 minutes, well past + # spur's ~1-2 min registration lag. + _SPUR_STARTUP_POLLS = 20 + + def monitor(self, deployment_id: str) -> DeploymentResult: + """Marker-based completion detection for the spur job array. + + Each array task writes ``done_rank`` (its exit code) into + ``//`` on the shared filesystem. We treat + those markers as the source of truth because spur's ``sacct -j`` does not + filter by job id and ``squeue`` is eventually consistent. + """ + marker_dir = Path(self.rendezvous_dir) / str(deployment_id) + n = int(self.nodes) + live_output = self.config.additional_context.get("live_output", False) + + codes: Dict[int, int] = {} + if marker_dir.is_dir(): + for rank in range(n): + f = marker_dir / f"done_rank{rank}" + if f.exists(): + try: + codes[rank] = int((f.read_text().strip() or "1")) + except ValueError: + codes[rank] = 1 + + if len(codes) >= n: + failed = {r: c for r, c in codes.items() if c != 0} + self._report_logs( + deployment_id, success=not failed, live_output=live_output + ) + if not failed: + return DeploymentResult( + status=DeploymentStatus.SUCCESS, + deployment_id=deployment_id, + message=f"All {n} array tasks completed successfully", + ) + return DeploymentResult( + status=DeploymentStatus.FAILED, + deployment_id=deployment_id, + message=f"Array task(s) failed (rank:exit) {failed}", + ) + + # Not all ranks done yet. Guard against a task that died without writing a + # marker, but only AFTER we have seen the tasks alive at least once: right + # after sbatch, spur's squeue does not yet list the array tasks, so a fresh + # healthy run legitimately reports 0 live tasks for the first ~1-2 min. + live = self._live_task_count(deployment_id, self._model_job_name()) + if live > 0: + self._spur_seen_live = True + self._spur_empty_polls = 0 + self._spur_unknown_polls = 0 + self._spur_startup_polls = 0 + elif live == 0 and getattr(self, "_spur_seen_live", False) and len(codes) < n: + # Tasks were running earlier and now none are queued and not all + # ranks reported: a transient empty squeue is possible, so require + # several consecutive empty polls before declaring failure. + self._spur_unknown_polls = 0 + self._spur_empty_polls = getattr(self, "_spur_empty_polls", 0) + 1 + if self._spur_empty_polls >= self._SPUR_DEAD_POLLS: + self._report_logs(deployment_id, success=False, live_output=live_output) + return DeploymentResult( + status=DeploymentStatus.FAILED, + deployment_id=deployment_id, + message=( + f"Only {len(codes)}/{n} ranks reported completion and no " + f"array tasks remain in the queue" + ), + ) + elif live < 0: + # squeue could not be queried. Bound this too: otherwise a control + # plane that stays down leaves monitor() returning RUNNING forever + # (the caller polls without a timeout). + self._spur_empty_polls = 0 + self._spur_unknown_polls = getattr(self, "_spur_unknown_polls", 0) + 1 + if self._spur_unknown_polls >= self._SPUR_UNKNOWN_POLLS: + self._report_logs(deployment_id, success=False, live_output=live_output) + return DeploymentResult( + status=DeploymentStatus.UNKNOWN, + deployment_id=deployment_id, + message=( + f"squeue unavailable for {self._spur_unknown_polls} consecutive " + f"polls and only {len(codes)}/{n} ranks reported completion" + ), + ) + else: + # Still in the startup grace window (tasks not yet registered), which + # is bounded so a job that dies before squeue ever lists it does not + # poll forever. + self._spur_empty_polls = 0 + self._spur_unknown_polls = 0 + self._spur_startup_polls = getattr(self, "_spur_startup_polls", 0) + 1 + if self._spur_startup_polls >= self._SPUR_STARTUP_POLLS: + self._report_logs(deployment_id, success=False, live_output=live_output) + return DeploymentResult( + status=DeploymentStatus.FAILED, + deployment_id=deployment_id, + message=( + f"No array task for job {deployment_id} was ever seen in the " + f"queue over {self._spur_startup_polls} polls and only " + f"{len(codes)}/{n} ranks reported completion" + ), + ) + + if live_output: + self._stream_job_output(deployment_id) + + return DeploymentResult( + status=DeploymentStatus.RUNNING, + deployment_id=deployment_id, + message=f"{len(codes)}/{n} ranks done (live tasks: {live})", + ) + + def _report_logs( + self, deployment_id: str, success: bool, live_output: bool + ) -> None: + """Emit final logs the same way SlurmDeployment.monitor() does. + + Args: + deployment_id: Array job id. + success: Whether the run succeeded (only used for the summary). + live_output: Whether the user asked for streamed output. + """ + if live_output: + self._stream_job_output(deployment_id, final=True) + else: + self._show_log_summary(deployment_id, success=success) diff --git a/src/madengine/deployment/templates/slurm/job.sh.j2 b/src/madengine/deployment/templates/slurm/job.sh.j2 index c692acf3..e991e889 100644 --- a/src/madengine/deployment/templates/slurm/job.sh.j2 +++ b/src/madengine/deployment/templates/slurm/job.sh.j2 @@ -1,11 +1,29 @@ #!/bin/bash #SBATCH --job-name=madengine-{{ model_name }} +{% if scheduler == 'spur' %} +# Log names use %A (array job id, the id sbatch returns) rather than %j (each +# array task's own job id) so result collection can find them by deployment id. +# %a is the array index, i.e. the node rank. +#SBATCH --output={{ output_dir }}/madengine-{{ model_name }}_%A_%a.out +#SBATCH --error={{ output_dir }}/madengine-{{ model_name }}_%A_%a.err +{% else %} #SBATCH --output={{ output_dir }}/madengine-{{ model_name }}_%j_%t.out #SBATCH --error={{ output_dir }}/madengine-{{ model_name }}_%j_%t.err +{% endif %} #SBATCH --partition={{ partition }} +{% if scheduler == 'spur' %} +# Spur scheduler: srun cannot fan out tasks across nodes, so instead of one +# multi-node job with `srun bash task` we use a job ARRAY of single-node tasks. +# Each array task runs on one node; SLURM_ARRAY_TASK_ID is the node rank, and +# the tasks self-form the cluster via the model launcher's TCP rendezvous. +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --array=0-{{ nodes - 1 }} +{% else %} #SBATCH --nodes={{ nodes }} #SBATCH --ntasks={{ nodes }} #SBATCH --ntasks-per-node=1 +{% endif %} {% if not skip_gpus_directive %}#SBATCH --gpus-per-node={{ gpus_per_node }} {% endif %} #SBATCH --time={{ time_limit }} @@ -60,7 +78,59 @@ export PATH="{{ submission_bin_dir }}:$PATH" # ============================================================================= # Distributed execution environment (auto-configured from SLURM) -export MASTER_ADDR=$(scontrol show hostname $SLURM_NODELIST | head -n 1) +# Scheduler-portable nodelist expansion. +# Stock SLURM: SLURM_NODELIST is compressed ("node[01-03]") -> needs +# `scontrol show hostnames`. Spur (Crusoe): SLURM_NODELIST is already an +# expanded comma list ("nodeA,nodeB") and `scontrol show hostname[s]` is +# unsupported. Try scontrol only for compressed forms, else split on commas. +mad_expand_nodelist() { + local nl="$1" + if [ -z "$nl" ]; then return 0; fi + if [[ "$nl" != *"["* ]]; then + echo "$nl" | tr ',' '\n' + return 0 + fi + if command -v scontrol >/dev/null 2>&1 && scontrol show hostnames "$nl" >/dev/null 2>&1; then + scontrol show hostnames "$nl" + else + echo "$nl" | tr ',' '\n' + fi +} +# Export so the generated TASK_SCRIPT (run as a separate `bash` process) can use +# it. Callers still guard with `command -v` in case the scheduler strips +# BASH_FUNC_* entries from the task environment. +export -f mad_expand_nodelist +{% if scheduler == 'spur' %} +# --- Spur job-array rank + rendezvous ---------------------------------------- +# Each array task is one node. SLURM_ARRAY_TASK_ID is the node rank. We emulate +# the SLURM_* vars the rest of this script expects, then resolve MASTER_ADDR via +# a shared-filesystem rendezvous: rank 0 publishes its transport IP; peers wait +# for and read it. The port is derived from the shared SLURM_ARRAY_JOB_ID so it +# is identical across all tasks (mirrors the model launcher's rendezvous). +export SLURM_PROCID="${SLURM_ARRAY_TASK_ID:-0}" +# Each array task is its own single-node allocation, so SLURM_NODEID is 0 +# everywhere; MAD_NODE_RANK (exported below for model scripts) needs the rank. +export SLURM_NODEID="${SLURM_ARRAY_TASK_ID:-0}" +export SLURM_NNODES={{ nodes }} +export SLURM_NTASKS={{ nodes }} +export SLURM_LOCALID="${SLURM_LOCALID:-0}" +# The model launcher (run.sh) derives its cross-node rendezvous PORT and its +# shared /run_logs/ readiness dir from SLURM_JOB_ID, which therefore MUST be +# identical on every node. Under a job array each task has its own SLURM_JOB_ID +# but shares SLURM_ARRAY_JOB_ID, so pin SLURM_JOB_ID to the array id for the +# container/launcher. (The array task id is preserved via SLURM_PROCID above and +# in the per-node TASK_SCRIPT filename below.) +export SLURM_JOB_ID="${SLURM_ARRAY_JOB_ID:-$SLURM_JOB_ID}" +# spur may not set SLURM_SUBMIT_DIR; the manifest mounts it for /run_logs. +export SLURM_SUBMIT_DIR="${SLURM_SUBMIT_DIR:-{{ manifest_file | dirname }}}" +{{ spur_rendezvous_block }} +export MASTER_PORT={{ master_port | default(29500) }} +export WORLD_SIZE={{ nodes }} +export LOCAL_RANK=0 +export NNODES={{ nodes }} +export GPUS_PER_NODE={{ gpus_per_node }} +{% else %} +export MASTER_ADDR=$(mad_expand_nodelist "$SLURM_NODELIST" | head -n 1) export MASTER_PORT={{ master_port | default(29500) }} export WORLD_SIZE=$SLURM_NTASKS # NOTE: RANK is set per-task inside srun context (not here in main script) @@ -68,6 +138,7 @@ export WORLD_SIZE=$SLURM_NTASKS export LOCAL_RANK=$SLURM_LOCALID export NNODES={{ nodes }} export GPUS_PER_NODE={{ gpus_per_node }} +{% endif %} # GPU visibility (ROCm/CUDA) # IMPORTANT: Ray (vLLM, SGLang) requires HIP_VISIBLE_DEVICES for AMD GPUs @@ -125,6 +196,11 @@ export MAD_DEPLOYMENT_TYPE=slurm export MAD_SLURM_JOB_ID=$SLURM_JOB_ID export MAD_NODE_RANK=$SLURM_NODEID export MAD_TOTAL_NODES={{ nodes }} +# Job id used for cross-node artifact/log collection. On stock SLURM this is the +# job id. On spur every array task has its own SLURM_JOB_ID but shares +# SLURM_ARRAY_JOB_ID, which is what collect_results() keys on (== deploy() id). +# Exported so srun/bash task children inherit the same value. +export MAD_COLLECT_JOB_ID="${SLURM_ARRAY_JOB_ID:-$SLURM_JOB_ID}" # ============================================================================= # Workspace Setup @@ -403,7 +479,14 @@ export TORCH_ELASTIC_RDZV_TIMEOUT=3600 # Use submission directory (shared filesystem) for task script # /tmp is local to each node and won't be accessible by srun on other nodes +{% if scheduler == 'spur' %} +# Under a job array SLURM_JOB_ID is pinned to the shared array id, so include the +# node rank (array task id) to keep each node's task script path unique on the +# shared filesystem. +TASK_SCRIPT="{{ manifest_file | dirname }}/{{ output_dir }}/madengine_task_${SLURM_JOB_ID}_${SLURM_PROCID}.sh" +{% else %} TASK_SCRIPT="{{ manifest_file | dirname }}/{{ output_dir }}/madengine_task_${SLURM_JOB_ID}.sh" +{% endif %} cat > "$TASK_SCRIPT" << 'TASK_SCRIPT_EOF' #!/bin/bash @@ -463,7 +546,7 @@ echo "" if [ -n "$SLURM_TMPDIR" ] && [ -d "$SLURM_TMPDIR" ] && [ -w "$SLURM_TMPDIR" ]; then WORKSPACE=$SLURM_TMPDIR/madengine_node_${SLURM_PROCID} else - WORKSPACE=/tmp/madengine_job_${SLURM_JOB_ID}_node_${SLURM_PROCID} + WORKSPACE=/tmp/madengine_job_${MAD_COLLECT_JOB_ID}_node_${SLURM_PROCID} fi mkdir -p $WORKSPACE WORKSPACE_TYPE="local-multinode" @@ -498,8 +581,8 @@ echo "" # Debug: set node log paths early and trap EXIT so we always see which node failed # ============================================================================= RESULTS_DIR="$SUBMISSION_DIR" -NODE_LOG_OUT="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${SLURM_PROCID}.out" -NODE_LOG_ERR="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${SLURM_PROCID}.err" +NODE_LOG_OUT="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${SLURM_PROCID}.out" +NODE_LOG_ERR="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${SLURM_PROCID}.err" echo "[DEBUG] $(date -Iseconds) Node ${SLURM_PROCID} ($(hostname)): task script started" >> "${NODE_LOG_ERR}" trap 'ec=$?; echo "[DEBUG] $(date -Iseconds) Node ${SLURM_PROCID} ($(hostname)): task script exiting with code $ec" >> "${NODE_LOG_ERR}"' EXIT @@ -654,8 +737,8 @@ echo "" # Create node-specific log files in results directory RESULTS_DIR={{ manifest_file | dirname }} -NODE_LOG_OUT="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${SLURM_PROCID}.out" -NODE_LOG_ERR="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${SLURM_PROCID}.err" +NODE_LOG_OUT="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${SLURM_PROCID}.out" +NODE_LOG_ERR="${RESULTS_DIR}/{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${SLURM_PROCID}.err" echo "Node ${SLURM_PROCID} logs:" echo " stdout: ${NODE_LOG_OUT}" @@ -696,7 +779,7 @@ echo "Task completed with exit code: $TASK_EXIT" # ============================================================================= SUBMISSION_DIR={{ manifest_file | dirname }} -JOB_COLLECTION_DIR="${SUBMISSION_DIR}/{{ output_dir }}/{{ model_name }}/${SLURM_JOB_ID}" +JOB_COLLECTION_DIR="${SUBMISSION_DIR}/{{ output_dir }}/{{ model_name }}/${MAD_COLLECT_JOB_ID}" NODE_COLLECTION_DIR="${JOB_COLLECTION_DIR}/node_${SLURM_PROCID}" mkdir -p "$NODE_COLLECTION_DIR" @@ -735,9 +818,18 @@ TASK_SCRIPT_EOF chmod +x "$TASK_SCRIPT" +{% if scheduler == 'spur' %} +# Spur: this script IS a single array task pinned to one node. Run the per-node +# task body directly (spur srun cannot dispatch to other nodes anyway). Node +# fan-out is achieved by the job array; peers coordinate via the rendezvous above. +echo "Launching task for node rank ${SLURM_PROCID} (array) on $(hostname)..." +bash "$TASK_SCRIPT" +EXIT_CODE=$? +{% else %} echo "Launching tasks on {{ nodes }} nodes..." srun bash "$TASK_SCRIPT" EXIT_CODE=$? +{% endif %} # Cleanup task script rm -f "$TASK_SCRIPT" @@ -776,7 +868,7 @@ EXIT_CODE=$? {% if multiple_results %} if [ $EXIT_CODE -eq 0 ]; then SUBMISSION_DIR={{ manifest_file | dirname }} - JOB_COLLECTION_DIR="${SUBMISSION_DIR}/{{ output_dir }}/{{ model_name }}/${SLURM_JOB_ID}" + JOB_COLLECTION_DIR="${SUBMISSION_DIR}/{{ output_dir }}/{{ model_name }}/${MAD_COLLECT_JOB_ID}" NODE_COLLECTION_DIR="${JOB_COLLECTION_DIR}/node_0" mkdir -p "$NODE_COLLECTION_DIR" if [ -f "$WORKSPACE/run_directory/{{ multiple_results }}" ]; then cp "$WORKSPACE/run_directory/{{ multiple_results }}" "$NODE_COLLECTION_DIR/" 2>/dev/null || true; fi @@ -785,6 +877,16 @@ fi {% endif %} {% endif %} +{% if scheduler == 'spur' %} +# Spur: publish this array task's exit code so the submitter's monitor() can +# detect per-rank completion on the shared filesystem. spur `sacct -j` does not +# filter by job id (returns the whole cluster history), so state cannot be read +# reliably per job; these markers are the source of truth for completion. +_MAD_DONE_DIR="{{ rendezvous_dir }}/${SLURM_ARRAY_JOB_ID:-$SLURM_JOB_ID}" +mkdir -p "$_MAD_DONE_DIR" 2>/dev/null || true +echo "$EXIT_CODE" > "$_MAD_DONE_DIR/done_rank${SLURM_PROCID}" +{% endif %} + # ============================================================================= # Job Completion # ============================================================================= @@ -808,8 +910,8 @@ if [ $EXIT_CODE -eq 0 ]; then echo " 📋 Individual Node Logs ({{ nodes }} nodes):" echo " ─────────────────────────────────────────────" for i in $(seq 0 $(({{ nodes }} - 1))); do - NODE_OUT="{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${i}.out" - NODE_ERR="{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${i}.err" + NODE_OUT="{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${i}.out" + NODE_ERR="{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${i}.err" if [ -f "$NODE_OUT" ]; then OUT_SIZE=$(du -h "$NODE_OUT" 2>/dev/null | cut -f1) ERR_SIZE=$(du -h "$NODE_ERR" 2>/dev/null | cut -f1) @@ -831,14 +933,14 @@ else echo " 📋 Check Individual Node Logs:" echo " ─────────────────────────────────" for i in $(seq 0 $(({{ nodes }} - 1))); do - NODE_OUT="{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${i}.out" - NODE_ERR="{{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_node_${i}.err" + NODE_OUT="{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${i}.out" + NODE_ERR="{{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_node_${i}.err" if [ -f "$NODE_OUT" ] || [ -f "$NODE_ERR" ]; then echo " Node $i: ${NODE_OUT}" fi done {% else %} - echo " Check logs: {{ output_dir }}/madengine-{{ model_name }}_${SLURM_JOB_ID}_*.out" + echo " Check logs: {{ output_dir }}/madengine-{{ model_name }}_${MAD_COLLECT_JOB_ID}_*.out" {% endif %} echo "========================================================================" fi diff --git a/src/madengine/orchestration/build_orchestrator.py b/src/madengine/orchestration/build_orchestrator.py index 48428b89..ec105a19 100644 --- a/src/madengine/orchestration/build_orchestrator.py +++ b/src/madengine/orchestration/build_orchestrator.py @@ -1319,7 +1319,11 @@ def _save_deployment_config(self, manifest_file: str): if not target: # Auto-detect based on config presence if self.additional_context.get("slurm"): - target = "slurm" + # Spur reuses the "slurm" block; slurm.scheduler picks the flavor. + scheduler = str( + self.additional_context["slurm"].get("scheduler", "") or "" + ).lower() + target = "spur" if scheduler == "spur" else "slurm" elif self.additional_context.get("k8s") or self.additional_context.get("kubernetes"): target = "k8s" else: diff --git a/src/madengine/orchestration/run_orchestrator.py b/src/madengine/orchestration/run_orchestrator.py index b0f0cb55..96515533 100644 --- a/src/madengine/orchestration/run_orchestrator.py +++ b/src/madengine/orchestration/run_orchestrator.py @@ -21,22 +21,24 @@ from rich.console import Console as RichConsole from rich.panel import Panel -from madengine.core.console import Console from madengine.core.auth import load_credentials +from madengine.core.console import Console from madengine.core.context import Context from madengine.core.dataprovider import Data -from madengine.core.timeout import resolve_run_timeout from madengine.core.errors import ( BuildError, ConfigurationError, ExecutionError, create_error_context, ) -from madengine.utils.session_tracker import SessionTracker +from madengine.core.timeout import resolve_run_timeout from madengine.orchestration.image_filtering import ( filter_images_by_gpu_compatibility as _filter_by_gpu_compat, +) +from madengine.orchestration.image_filtering import ( filter_images_by_skip_gpu_arch as _filter_by_skip_gpu_arch, ) +from madengine.utils.session_tracker import SessionTracker class RunOrchestrator: @@ -70,14 +72,19 @@ def __init__(self, args, additional_context: Optional[Dict] = None): # Use ast.literal_eval for Python dict syntax (single quotes) # This matches what Context class expects import ast + parsed = ast.literal_eval(args.additional_context) merged_context = parsed if isinstance(parsed, dict) else {} elif isinstance(args.additional_context, dict): merged_context = args.additional_context except (ValueError, SyntaxError) as e: - self.rich_console.print(f"[yellow]Warning: Could not parse additional_context: {e}[/yellow]") + self.rich_console.print( + f"[yellow]Warning: Could not parse additional_context: {e}[/yellow]" + ) if args.additional_context: - self.rich_console.print(f"[dim]Raw (first 200 chars): {str(args.additional_context)[:200]}[/dim]") + self.rich_console.print( + f"[dim]Raw (first 200 chars): {str(args.additional_context)[:200]}[/dim]" + ) pass if additional_context: @@ -91,19 +98,25 @@ def __init__(self, args, additional_context: Optional[Dict] = None): if getattr(args, "require_pinned_image", False): self.additional_context["require_pinned_image"] = True - keys_str = ", ".join(sorted(self.additional_context.keys())) if self.additional_context else "(none)" - self.rich_console.print(f"[dim]Run additional context (CLI):[/dim] [cyan]{keys_str}[/cyan]") + keys_str = ( + ", ".join(sorted(self.additional_context.keys())) + if self.additional_context + else "(none)" + ) + self.rich_console.print( + f"[dim]Run additional context (CLI):[/dim] [cyan]{keys_str}[/cyan]" + ) # Track if we copied MODEL_DIR contents (for cleanup) self._copied_from_model_dir = False - + # Track if we ran build phase in this workflow (for log combination) self._did_build_phase = False - + # Initialize session tracker for filtering current run results perf_csv_path = getattr(args, "output", "perf.csv") self.session_tracker = SessionTracker(perf_csv_path) - + # Initialize context in runtime mode (with GPU detection for local) # This will be lazy-initialized only when needed self.context = None @@ -113,14 +126,14 @@ def _init_runtime_context(self): """Initialize runtime context (with GPU detection).""" # Always reinitialize context in runtime mode for run phase # This ensures GPU detection and proper runtime context even after build phase - + # Context expects additional_context as a string representation of Python dict # Use repr() instead of json.dumps() because Context uses ast.literal_eval() if self.additional_context: context_string = repr(self.additional_context) else: context_string = None - + self.context = Context( additional_context=context_string, build_only_mode=False, @@ -181,7 +194,7 @@ def execute( mad_container_image = None if self.additional_context: mad_container_image = self.additional_context.get("MAD_CONTAINER_IMAGE") - + if mad_container_image: # Local image mode: Skip build, create synthetic manifest if not tags: @@ -196,14 +209,16 @@ def execute( "Example: --tags model_name --additional-context \"{'MAD_CONTAINER_IMAGE': 'rocm/tensorflow:latest'}\"", ], ) - + # Generate synthetic manifest using the provided image manifest_file = self._create_manifest_from_local_image( image_name=mad_container_image, tags=tags, - manifest_output=getattr(self.args, "manifest_output", "build_manifest.json"), + manifest_output=getattr( + self.args, "manifest_output", "build_manifest.json" + ), ) - + # Step 1: Ensure we have a manifest (build if needed) elif not manifest_file or not os.path.exists(manifest_file): if not tags: @@ -219,7 +234,9 @@ def execute( ], ) - self.rich_console.print("[cyan]No manifest found, building first...[/cyan]\n") + self.rich_console.print( + "[cyan]No manifest found, building first...[/cyan]\n" + ) manifest_file = self._build_phase(tags, registry) self._did_build_phase = True # Mark that we built in this workflow @@ -230,44 +247,66 @@ def execute( # (with optional runtime override) with open(manifest_file) as f: manifest = json.load(f) - + deployment_config = manifest.get("deployment_config", {}) - + # Update additional_context with deployment_config for deployment layer if not self.additional_context: self.additional_context = {} - + # Merge deployment_config into additional_context (for deployment layer to use) - for key in ["slurm", "k8s", "kubernetes", "distributed", "vllm", "env_vars", "debug"]: + for key in [ + "slurm", + "k8s", + "kubernetes", + "distributed", + "vllm", + "env_vars", + "debug", + ]: if key in deployment_config and key not in self.additional_context: self.additional_context[key] = deployment_config[key] - + # Display manifest entries: context (from build) and deployment_config (run/deploy) self.rich_console.print("[bold blue]Build manifest breakdown[/bold blue]\n") manifest_context = manifest.get("context", {}) - self.rich_console.print(Panel( - json.dumps(manifest_context, indent=2) if manifest_context else "(empty)", - title="[bold]Manifest context[/bold] (from build additional context)", - border_style="dim", - padding=(0, 1), - )) - self.rich_console.print(Panel( - json.dumps(deployment_config, indent=2) if deployment_config else "(empty)", - title="[bold]Manifest deployment_config[/bold]", - border_style="dim", - padding=(0, 1), - )) + self.rich_console.print( + Panel( + ( + json.dumps(manifest_context, indent=2) + if manifest_context + else "(empty)" + ), + title="[bold]Manifest context[/bold] (from build additional context)", + border_style="dim", + padding=(0, 1), + ) + ) + self.rich_console.print( + Panel( + ( + json.dumps(deployment_config, indent=2) + if deployment_config + else "(empty)" + ), + title="[bold]Manifest deployment_config[/bold]", + border_style="dim", + padding=(0, 1), + ) + ) self.rich_console.print() # Infer deployment target from config structure (Convention over Configuration) # No explicit "deploy" field needed - presence of k8s/slurm indicates deployment type target = self._infer_deployment_target(self.additional_context) - + # Legacy support: check manifest for explicit target if not target or target == "local": target = deployment_config.get("target", "local") - - self.rich_console.print(f"[bold cyan]Deployment target: {target}[/bold cyan]\n") + + self.rich_console.print( + f"[bold cyan]Deployment target: {target}[/bold cyan]\n" + ) # Step 4: Execute based on target try: @@ -275,28 +314,34 @@ def execute( results = self._execute_local(manifest_file, timeout) else: results = self._execute_distributed(target, manifest_file) - + # Combine build and run logs for full workflow if self._did_build_phase and (target == "local" or target == "docker"): self._combine_build_and_run_logs(manifest_file) - + # Add session information to results for filtering results["session_start_row"] = session_start_row - results["session_row_count"] = self.session_tracker.get_session_row_count() - + results["session_row_count"] = ( + self.session_tracker.get_session_row_count() + ) + # Always cleanup madengine package files after execution - self.rich_console.print("\n[dim]🧹 Cleaning up madengine package files...[/dim]") + self.rich_console.print( + "\n[dim]🧹 Cleaning up madengine package files...[/dim]" + ) self._cleanup_model_dir_copies() - + # NOTE: Do NOT cleanup session marker here! # It's needed by display functions in CLI layer # Cleanup happens in CLI after display (via perf_csv_path) - + return results - + except Exception as e: # Always cleanup madengine package files even on error - self.rich_console.print("\n[dim]🧹 Cleaning up madengine package files...[/dim]") + self.rich_console.print( + "\n[dim]🧹 Cleaning up madengine package files...[/dim]" + ) self._cleanup_model_dir_copies() raise @@ -342,59 +387,68 @@ def _build_phase(self, tags: list, registry: Optional[str] = None) -> str: return manifest_file def _create_manifest_from_local_image( - self, - image_name: str, - tags: list, - manifest_output: str = "build_manifest.json" + self, image_name: str, tags: list, manifest_output: str = "build_manifest.json" ) -> str: """ Create a synthetic manifest for a user-provided local image. - + This enables MAD_CONTAINER_IMAGE functionality where users can skip the build phase and directly run models using a pre-existing Docker image. - + Args: image_name: Docker image name/tag (e.g., 'rocm/tensorflow:latest') tags: Model tags to discover manifest_output: Output path for the manifest file - + Returns: Path to the generated manifest file - + Raises: DiscoveryError: If no models are found RuntimeError: If image validation fails """ - from madengine.utils.discover_models import DiscoverModels from madengine.core.errors import DiscoveryError - - self.rich_console.print(f"[yellow]🏠 Local Image Mode: Using {image_name}[/yellow]") - self.rich_console.print(f"[dim]Skipping build phase, creating synthetic manifest...[/dim]\n") - + from madengine.utils.discover_models import DiscoverModels + + self.rich_console.print( + f"[yellow]🏠 Local Image Mode: Using {image_name}[/yellow]" + ) + self.rich_console.print( + f"[dim]Skipping build phase, creating synthetic manifest...[/dim]\n" + ) + # Validate that the image exists locally or can be pulled. # image_name is interpolated into shell commands run with shell=True, # so shell-escape it to avoid command injection / breakage on special chars. quoted_image_name = shlex.quote(image_name) try: - self.console.sh(f"docker image inspect {quoted_image_name} > /dev/null 2>&1") - self.rich_console.print(f"[green]✓ Image {image_name} found locally[/green]") + self.console.sh( + f"docker image inspect {quoted_image_name} > /dev/null 2>&1" + ) + self.rich_console.print( + f"[green]✓ Image {image_name} found locally[/green]" + ) except (subprocess.CalledProcessError, RuntimeError) as e: - self.rich_console.print(f"[yellow]⚠️ Image {image_name} not found locally, attempting to pull...[/yellow]") + self.rich_console.print( + f"[yellow]⚠️ Image {image_name} not found locally, attempting to pull...[/yellow]" + ) try: self.console.sh(f"docker pull {quoted_image_name}") - self.rich_console.print(f"[green]✓ Successfully pulled {image_name}[/green]") + self.rich_console.print( + f"[green]✓ Successfully pulled {image_name}[/green]" + ) except Exception as e: raise RuntimeError( f"Failed to find or pull image {image_name}. " f"Ensure the image exists locally or can be pulled from a registry. " f"Error: {e}" ) - + # Discover models by tags (without building) self.args.tags = tags discover_models = DiscoverModels(args=self.args) models = discover_models.run() - + if not models: raise DiscoveryError( "No models discovered for local image mode", @@ -408,17 +462,21 @@ def _create_manifest_from_local_image( "Ensure model definitions have matching tags", ], ) - - self.rich_console.print(f"[green]✓ Discovered {len(models)} model(s) for tags: {tags}[/green]\n") - + + self.rich_console.print( + f"[green]✓ Discovered {len(models)} model(s) for tags: {tags}[/green]\n" + ) + # Initialize build-only context for manifest generation # (we need context structure, but skip GPU detection since we're not building) - context_string = repr(self.additional_context) if self.additional_context else None + context_string = ( + repr(self.additional_context) if self.additional_context else None + ) build_context = Context( additional_context=context_string, build_only_mode=True, ) - + # Create manifest structure manifest = { "built_images": {}, @@ -428,13 +486,13 @@ def _create_manifest_from_local_image( "local_image_name": image_name, "deployment_config": self.additional_context.get("deployment_config", {}), } - + # For each model, create a synthetic entry using the provided image for model in models: model_name = model["name"] # Create a synthetic image identifier (not an actual built image) synthetic_image_id = f"local-{model_name.replace('/', '_')}" - + manifest["built_images"][synthetic_image_id] = { "docker_image": image_name, # Use user-provided image "dockerfile": "N/A (local image mode)", @@ -443,22 +501,26 @@ def _create_manifest_from_local_image( "local_image": True, "registry_image": None, } - + # Convert data list to comma-separated string (required by dataprovider) data_field = model.get("data", []) if isinstance(data_field, list): data_str = ",".join(data_field) if data_field else "" else: data_str = data_field if data_field else "" - + # Build model info dict with all fields that ContainerRunner expects # Use exact field names from models.json format manifest["built_models"][synthetic_image_id] = { "name": model_name, "tags": model.get("tags", []), "dockerfile": "N/A (local image mode)", - "scripts": model.get("scripts", ""), # models.json uses "scripts" (plural) - "n_gpus": model.get("n_gpus", "1"), # models.json uses "n_gpus" (string format) + "scripts": model.get( + "scripts", "" + ), # models.json uses "scripts" (plural) + "n_gpus": model.get( + "n_gpus", "1" + ), # models.json uses "n_gpus" (string format) "owner": model.get("owner", ""), "training_precision": model.get("training_precision", ""), "args": model.get("args", ""), # Required field for docker run @@ -469,17 +531,23 @@ def _create_manifest_from_local_image( "cred": model.get("cred", ""), "deprecated": model.get("deprecated", False), "skip_gpu_arch": model.get("skip_gpu_arch", []), - "additional_docker_run_options": model.get("additional_docker_run_options", ""), + "additional_docker_run_options": model.get( + "additional_docker_run_options", "" + ), "multiple_results": model.get("multiple_results", ""), } - + # Write manifest to file with open(manifest_output, "w") as f: json.dump(manifest, f, indent=2) - - self.rich_console.print(f"[green]✓ Generated synthetic manifest: {manifest_output}[/green]") - self.rich_console.print(f"[yellow]⚠️ Warning: User-provided image {image_name}. Model support not guaranteed.[/yellow]\n") - + + self.rich_console.print( + f"[green]✓ Generated synthetic manifest: {manifest_output}[/green]" + ) + self.rich_console.print( + f"[yellow]⚠️ Warning: User-provided image {image_name}. Model support not guaranteed.[/yellow]\n" + ) + return manifest_output def _load_and_merge_manifest(self, manifest_file: str) -> str: @@ -498,15 +566,24 @@ def _load_and_merge_manifest(self, manifest_file: str) -> str: if "deployment_config" in manifest: stored_config = manifest["deployment_config"] # Runtime --additional-context overrides stored config - for key in ["deploy", "slurm", "k8s", "kubernetes", "distributed", "vllm", "env_vars", "debug"]: + for key in [ + "deploy", + "slurm", + "k8s", + "kubernetes", + "distributed", + "vllm", + "env_vars", + "debug", + ]: if key in self.additional_context: stored_config[key] = self.additional_context[key] manifest["deployment_config"] = stored_config - + # Merge context (tools, pre_scripts, post_scripts, encapsulate_script) if "context" not in manifest: manifest["context"] = {} - + merge_keys = [ "tools", "pre_scripts", @@ -521,7 +598,7 @@ def _load_and_merge_manifest(self, manifest_file: str) -> str: if key in self.additional_context: manifest["context"][key] = self.additional_context[key] context_updated = True - + if context_updated or "deployment_config" in manifest: # Write back merged config with open(manifest_file, "w") as f: @@ -537,16 +614,18 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: # Load manifest first to check if we have Docker images with open(manifest_file, "r") as f: manifest = json.load(f) - + has_docker_images = bool(manifest.get("built_images", {})) - + if has_docker_images: # Using Docker containers - containers have GPU support built-in - self.rich_console.print("[dim cyan]Using Docker containers with built-in GPU support[/dim cyan]\n") - + self.rich_console.print( + "[dim cyan]Using Docker containers with built-in GPU support[/dim cyan]\n" + ) + # Initialize runtime context (runs full GPU detection on compute nodes) self._init_runtime_context() - + # Show node info self._show_node_info() @@ -565,7 +644,9 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: if "docker_mounts" in manifest_context: if "docker_mounts" not in self.context.ctx: self.context.ctx["docker_mounts"] = {} - for container_path, host_path in manifest_context["docker_mounts"].items(): + for container_path, host_path in manifest_context[ + "docker_mounts" + ].items(): if container_path not in self.context.ctx["docker_mounts"]: self.context.ctx["docker_mounts"][container_path] = host_path if "docker_build_arg" in manifest_context: @@ -574,9 +655,15 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: for key, value in manifest_context["docker_build_arg"].items(): if key not in self.context.ctx["docker_build_arg"]: self.context.ctx["docker_build_arg"][key] = value - if "docker_gpus" in manifest_context and "docker_gpus" not in self.context.ctx: + if ( + "docker_gpus" in manifest_context + and "docker_gpus" not in self.context.ctx + ): self.context.ctx["docker_gpus"] = manifest_context["docker_gpus"] - if "gpu_vendor" in manifest_context and "gpu_vendor" not in self.context.ctx: + if ( + "gpu_vendor" in manifest_context + and "gpu_vendor" not in self.context.ctx + ): self.context.ctx["gpu_vendor"] = manifest_context["gpu_vendor"] if "guest_os" in manifest_context and "guest_os" not in self.context.ctx: self.context.ctx["guest_os"] = manifest_context["guest_os"] @@ -587,12 +674,17 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: if "post_scripts" in manifest_context: self.context.ctx["post_scripts"] = manifest_context["post_scripts"] if "encapsulate_script" in manifest_context: - self.context.ctx["encapsulate_script"] = manifest_context["encapsulate_script"] + self.context.ctx["encapsulate_script"] = manifest_context[ + "encapsulate_script" + ] # Restore docker_env_vars from build context (e.g. MAD_SECRETS_HFTOKEN for Primus HF-backed configs). # Keep runtime-detected values as priority (consistent with docker_mounts / docker_build_arg): # values already populated by Context (e.g. MAD_SECRETS_* read from os.environ) must not be # overwritten by manifest entries that may still contain unexpanded "${VAR}" placeholders. - if "docker_env_vars" in manifest_context and manifest_context["docker_env_vars"]: + if ( + "docker_env_vars" in manifest_context + and manifest_context["docker_env_vars"] + ): if "docker_env_vars" not in self.context.ctx: self.context.ctx["docker_env_vars"] = {} for k, v in manifest_context["docker_env_vars"].items(): @@ -610,9 +702,13 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: if "pre_scripts" in self.additional_context: self.context.ctx["pre_scripts"] = self.additional_context["pre_scripts"] if "post_scripts" in self.additional_context: - self.context.ctx["post_scripts"] = self.additional_context["post_scripts"] + self.context.ctx["post_scripts"] = self.additional_context[ + "post_scripts" + ] if "encapsulate_script" in self.additional_context: - self.context.ctx["encapsulate_script"] = self.additional_context["encapsulate_script"] + self.context.ctx["encapsulate_script"] = self.additional_context[ + "encapsulate_script" + ] # Filter images by GPU vendor and architecture # Filter images by GPU compatibility @@ -625,10 +721,14 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: if has_docker_images: # Docker images: filter by GPU vendor at runtime to avoid cross-vendor execution - self.rich_console.print("[dim cyan]Filtering Docker images by runtime GPU compatibility...[/dim cyan]") + self.rich_console.print( + "[dim cyan]Filtering Docker images by runtime GPU compatibility...[/dim cyan]" + ) else: # Bare-metal execution: filter by runtime GPU - self.rich_console.print("[dim cyan]Filtering bare-metal images by runtime GPU compatibility...[/dim cyan]") + self.rich_console.print( + "[dim cyan]Filtering bare-metal images by runtime GPU compatibility...[/dim cyan]" + ) compatible_images = self._filter_images_by_gpu_compatibility( manifest["built_images"], runtime_gpu_vendor, runtime_gpu_arch @@ -650,30 +750,37 @@ def _execute_local(self, manifest_file: str, timeout: int) -> Dict: manifest["built_images"] = compatible_images print(f"Filtered to {len(compatible_images)} compatible images\n") - + # Filter by skip_gpu_arch from model definitions (applies to both Docker and bare-metal) runtime_gpu_arch = self.context.get_system_gpu_architecture() if "built_models" in manifest and compatible_images: - self.rich_console.print("[cyan]Checking skip_gpu_arch model restrictions...[/cyan]") + self.rich_console.print( + "[cyan]Checking skip_gpu_arch model restrictions...[/cyan]" + ) compatible_images = self._filter_images_by_skip_gpu_arch( compatible_images, manifest["built_models"], runtime_gpu_arch ) manifest["built_images"] = compatible_images - print(f"After skip_gpu_arch filtering: {len(compatible_images)} images to run\n") - + print( + f"After skip_gpu_arch filtering: {len(compatible_images)} images to run\n" + ) + # NOTE: Dockerfile context filtering is already done during build phase # Re-filtering during run phase causes issues because: # 1. The build phase already filtered dockerfiles based on build-time context # 2. All built images should be runnable on the runtime node # 3. Legacy behavior: filtering happens once (either build or run, not both) - + # Write filtered manifest back to file so runner sees the filtered list with open(manifest_file, "w") as f: json.dump(manifest, f, indent=2) except Exception as e: import traceback - self.rich_console.print(f"[yellow]Warning: GPU/Context filtering failed: {e}[/yellow]") + + self.rich_console.print( + f"[yellow]Warning: GPU/Context filtering failed: {e}[/yellow]" + ) self.rich_console.print(f"[red]Traceback: {traceback.format_exc()}[/red]") self.rich_console.print("[yellow]Proceeding with all images[/yellow]\n") @@ -732,13 +839,15 @@ def _execute_distributed(self, target: str, manifest_file: str) -> Dict: ) # Import from deployment layer - from madengine.deployment.factory import DeploymentFactory from madengine.deployment.base import DeploymentConfig + from madengine.deployment.factory import DeploymentFactory # Add runtime flags to additional_context for deployment layer if "live_output" not in self.additional_context: - self.additional_context["live_output"] = getattr(self.args, "live_output", False) - + self.additional_context["live_output"] = getattr( + self.args, "live_output", False + ) + # Pass session_start_row for result filtering in collect_results session_start_row = self.session_tracker.session_start_row if "session_start_row" not in self.additional_context: @@ -758,9 +867,7 @@ def _execute_distributed(self, target: str, manifest_file: str) -> Dict: # itself, and a concrete value here would read as an explicit # --timeout and outrank the card. No model card is consulted at this # level, hence the empty dict. - timeout=resolve_run_timeout( - {}, getattr(self.args, "timeout", -1) - ), + timeout=resolve_run_timeout({}, getattr(self.args, "timeout", -1)), cli_timeout=getattr(self.args, "timeout", -1), monitor=self.additional_context.get("monitor", True), cleanup_on_failure=self.additional_context.get("cleanup_on_failure", True), @@ -808,37 +915,39 @@ def _show_node_info(self): elif "HOST_AZURE" in host_os: print(self.console.sh("timeout 10 tdnf info rocm-libs", canFail=True)) else: - self.rich_console.print("[yellow]Warning: Unable to detect host OS[/yellow]") + self.rich_console.print( + "[yellow]Warning: Unable to detect host OS[/yellow]" + ) def _cleanup_model_dir_copies(self): """Clean up only madengine package files from scripts/common directory. - + This cleanup removes ONLY the files that were copied from madengine package: - scripts/common/tools.json - scripts/common/test_echo.sh - scripts/common/pre_scripts/ - scripts/common/post_scripts/ - scripts/common/tools/ - + This preserves the user's actual scripts/ and docker/ directories in MAD project. """ import shutil import subprocess - + # Only clean up scripts/common/ subdirectories that came from madengine package common_dir = Path("scripts/common") if not common_dir.exists(): return - + # List of items to clean up (from madengine package) items_to_cleanup = [ "tools.json", "test_echo.sh", "pre_scripts", "post_scripts", - "tools" + "tools", ] - + for item_name in items_to_cleanup: item_path = common_dir / item_name if item_path.exists(): @@ -849,14 +958,20 @@ def _cleanup_model_dir_copies(self): subprocess.run( ["chmod", "-R", "+w", str(item_path)], capture_output=True, - timeout=10 + timeout=10, ) - except (subprocess.TimeoutExpired, subprocess.CalledProcessError, OSError) as e: + except ( + subprocess.TimeoutExpired, + subprocess.CalledProcessError, + OSError, + ) as e: print(f"Warning: chmod failed for {item_path}: {e}") shutil.rmtree(item_path) else: item_path.unlink() - self.rich_console.print(f"[dim] Cleaned up: scripts/common/{item_name}[/dim]") + self.rich_console.print( + f"[dim] Cleaned up: scripts/common/{item_name}[/dim]" + ) except Exception as e: # Try with sudo for permission issues try: @@ -864,9 +979,11 @@ def _cleanup_model_dir_copies(self): ["sudo", "rm", "-rf", str(item_path)], check=True, capture_output=True, - timeout=10 + timeout=10, + ) + self.rich_console.print( + f"[dim] Cleaned up: scripts/common/{item_name} (elevated)[/dim]" ) - self.rich_console.print(f"[dim] Cleaned up: scripts/common/{item_name} (elevated)[/dim]") except Exception as e2: self.rich_console.print( f"[yellow]⚠️ Warning: Could not clean up {item_path}: {e2}[/yellow]" @@ -874,84 +991,88 @@ def _cleanup_model_dir_copies(self): def _combine_build_and_run_logs(self, manifest_file: str): """Combine build.live.log and run.live.log into live.log for full workflow. - + For full workflow (build + run), this creates a unified log file by: 1. Reading the manifest to find models that were actually executed in this session 2. Finding corresponding *.build.live.log and *.run.live.log files for those models 3. Concatenating them into *.live.log 4. Keeping the original build and run logs for reference - + Args: manifest_file: Path to the manifest file containing executed models """ import json - + # Load manifest to get list of build log files try: with open(manifest_file, "r") as f: manifest = json.load(f) - + built_images = manifest.get("built_images", {}) if not built_images: return # No models to process except Exception as e: - self.rich_console.print(f"[yellow]⚠️ Warning: Could not load manifest for log combining: {e}[/yellow]") + self.rich_console.print( + f"[yellow]⚠️ Warning: Could not load manifest for log combining: {e}[/yellow]" + ) return - + self.rich_console.print("\n[dim]📝 Combining build and run logs...[/dim]") combined_count = 0 - + # Process each built image for image_name, image_info in built_images.items(): # Get build log file name from manifest build_log = image_info.get("log_file") if not build_log or not os.path.exists(build_log): continue # Skip if build log doesn't exist - + # Derive the base name and corresponding run log base_name = build_log.replace(".build.live.log", "") run_log = f"{base_name}.run.live.log" combined_log = f"{base_name}.live.log" - + # Check if run log exists if not os.path.exists(run_log): continue # Skip if run log doesn't exist - + try: # Combine build and run logs - with open(combined_log, 'w') as outfile: + with open(combined_log, "w") as outfile: # Add build log - with open(build_log, 'r') as infile: + with open(build_log, "r") as infile: outfile.write(infile.read()) - + # Add separator outfile.write("\n" + "=" * 80 + "\n") outfile.write("RUN PHASE LOG\n") outfile.write("=" * 80 + "\n\n") - + # Add run log - with open(run_log, 'r') as infile: + with open(run_log, "r") as infile: outfile.write(infile.read()) - + combined_count += 1 self.rich_console.print(f"[dim] Combined: {combined_log}[/dim]") - + except Exception as e: self.rich_console.print( f"[yellow]⚠️ Warning: Could not combine logs for {base_name}: {e}[/yellow]" ) - + if combined_count > 0: - self.rich_console.print(f"[dim]✓ Combined {combined_count} log file(s)[/dim]") + self.rich_console.print( + f"[dim]✓ Combined {combined_count} log file(s)[/dim]" + ) def _copy_scripts(self): """Copy common scripts to model directories. - + Handles scenarios: 1. MAD Project: scripts/ already exists in current directory - just add madengine common files 2. External MODEL_DIR: Copy from external path to current directory 3. madengine Testing: Copy from src/madengine/scripts/common - + NOTE: Does NOT delete existing scripts/ or docker/ directories in current working directory. """ import shutil @@ -959,19 +1080,27 @@ def _copy_scripts(self): # Define ignore function for cache files (used for all copy operations) def ignore_cache_files(directory, files): """Ignore Python cache files and directories.""" - return [f for f in files if f.endswith('.pyc') or f == '__pycache__' or f.endswith('.pyo')] - + return [ + f + for f in files + if f.endswith(".pyc") or f == "__pycache__" or f.endswith(".pyo") + ] + # Step 1: Check if MODEL_DIR points to external directory and copy if needed # MODEL_DIR default is "." (current directory), so only copy if it's different model_dir_env = os.environ.get("MODEL_DIR", ".") model_dir_abs = os.path.abspath(model_dir_env) current_dir_abs = os.path.abspath(".") - + # Only copy if MODEL_DIR points to a different directory (not current dir) if model_dir_abs != current_dir_abs and os.path.exists(model_dir_env): - self.rich_console.print(f"[yellow]📁 External MODEL_DIR detected: {model_dir_env}[/yellow]") - self.rich_console.print("[yellow]Copying MODEL_DIR contents for run phase...[/yellow]") - + self.rich_console.print( + f"[yellow]📁 External MODEL_DIR detected: {model_dir_env}[/yellow]" + ) + self.rich_console.print( + "[yellow]Copying MODEL_DIR contents for run phase...[/yellow]" + ) + # Copy docker/ and scripts/ from MODEL_DIR (without deleting existing ones first) for subdir in ["docker", "scripts"]: src_path = Path(model_dir_env) / subdir @@ -980,18 +1109,29 @@ def ignore_cache_files(directory, files): # Use copytree with dirs_exist_ok=True to merge instead of replace if dest_path.exists(): # Only warn, don't delete existing directories - self.rich_console.print(f"[dim] Note: Merging {subdir}/ from MODEL_DIR with existing directory[/dim]") - shutil.copytree(src_path, dest_path, dirs_exist_ok=True, ignore=ignore_cache_files) - - self.rich_console.print("[green]✓ MODEL_DIR structure copied (docker/, scripts/)[/green]") + self.rich_console.print( + f"[dim] Note: Merging {subdir}/ from MODEL_DIR with existing directory[/dim]" + ) + shutil.copytree( + src_path, + dest_path, + dirs_exist_ok=True, + ignore=ignore_cache_files, + ) + + self.rich_console.print( + "[green]✓ MODEL_DIR structure copied (docker/, scripts/)[/green]" + ) elif not os.path.exists(model_dir_env): - self.rich_console.print(f"[yellow]⚠️ Warning: MODEL_DIR '{model_dir_env}' does not exist, using current directory[/yellow]") + self.rich_console.print( + f"[yellow]⚠️ Warning: MODEL_DIR '{model_dir_env}' does not exist, using current directory[/yellow]" + ) # Step 2: Copy madengine's common scripts (pre_scripts, post_scripts, tools) # This provides the execution framework scripts # Find madengine installation path (works for both development and installed package) madengine_common = None - + # Option 1: Development mode - check if running from source dev_path = Path("src/madengine/scripts/common") if dev_path.exists(): @@ -1001,23 +1141,34 @@ def ignore_cache_files(directory, files): # Option 2: Installed package - find via module location try: import madengine + madengine_module_path = Path(madengine.__file__).parent installed_path = madengine_module_path / "scripts" / "common" if installed_path.exists(): madengine_common = installed_path - print(f"Found madengine scripts in installed package: {madengine_common}") + print( + f"Found madengine scripts in installed package: {madengine_common}" + ) except Exception as e: print(f"Could not locate madengine scripts: {e}") - + if madengine_common and madengine_common.exists(): - print(f"Copying madengine common scripts from {madengine_common} to scripts/common") - + print( + f"Copying madengine common scripts from {madengine_common} to scripts/common" + ) + dest_common = Path("scripts/common") # Ensure the destination directory exists before copying dest_common.mkdir(parents=True, exist_ok=True) - + # Copy pre_scripts, post_scripts, tools if they exist - for item in ["pre_scripts", "post_scripts", "tools", "tools.json", "test_echo.sh"]: + for item in [ + "pre_scripts", + "post_scripts", + "tools", + "tools.json", + "test_echo.sh", + ]: src_item = madengine_common / item if src_item.exists(): dest_item = dest_common / item @@ -1026,19 +1177,21 @@ def ignore_cache_files(directory, files): shutil.rmtree(dest_item) else: dest_item.unlink() - + if src_item.is_dir(): shutil.copytree(src_item, dest_item, ignore=ignore_cache_files) else: shutil.copy2(src_item, dest_item) print(f" Copied {item}") else: - self.rich_console.print("[yellow]⚠️ Could not find madengine scripts directory[/yellow]") + self.rich_console.print( + "[yellow]⚠️ Could not find madengine scripts directory[/yellow]" + ) # Step 3: REMOVED - Distribution to model directories is incorrect # scripts/common should remain at /scripts/common/ for proper relative path access # Model scripts reference it via ../scripts/common/ from their directory (e.g., scripts/dummy/) - # + # # This ensures compatibility with legacy workflow where: # - scripts/common/ stays at working directory root # - Model scripts use ../scripts/common/ relative paths @@ -1059,7 +1212,9 @@ def _filter_images_by_gpu_compatibility( ) compatible_images[model_name] = image_info continue - built_with_vendor = {k: v for k, v in built_images.items() if v.get("gpu_vendor")} + built_with_vendor = { + k: v for k, v in built_images.items() if v.get("gpu_vendor") + } compat, skipped = _filter_by_gpu_compat( built_with_vendor, runtime_gpu_vendor, runtime_gpu_arch ) @@ -1067,7 +1222,7 @@ def _filter_images_by_gpu_compatibility( for model_name, reason in skipped: self.rich_console.print(f"[dim] Skipping {model_name}: {reason}[/dim]") return compatible_images - + def _filter_images_by_gpu_architecture( self, built_images: Dict, runtime_gpu_arch: str ) -> Dict: @@ -1098,19 +1253,22 @@ def _filter_images_by_skip_gpu_arch( self._write_skipped_status(model_name, image_info, gpu_arch) return compatible_images - def _write_skipped_status(self, model_name: str, image_info: Dict, gpu_arch: str) -> None: + def _write_skipped_status( + self, model_name: str, image_info: Dict, gpu_arch: str + ) -> None: """Write SKIPPED status to perf CSV for models that were skipped. - + Args: model_name: Name of the model that was skipped image_info: Image information dictionary gpu_arch: GPU architecture that caused the skip """ try: - from madengine.reporting.update_perf_csv import update_perf_csv import json import tempfile - + + from madengine.reporting.update_perf_csv import update_perf_csv + # Create a perf entry for the skipped model perf_entry = { "model": model_name, @@ -1118,45 +1276,61 @@ def _write_skipped_status(self, model_name: str, image_info: Dict, gpu_arch: str "reason": f"Model not supported on {gpu_arch} architecture", "gpu_architecture": gpu_arch, } - + # Write to temporary JSON file - with tempfile.NamedTemporaryFile(mode='w', suffix='.json', delete=False) as f: + with tempfile.NamedTemporaryFile( + mode="w", suffix=".json", delete=False + ) as f: json.dump(perf_entry, f) temp_file = f.name - + # Get output CSV path from args - output_csv = getattr(self.args, 'output', 'perf.csv') - + output_csv = getattr(self.args, "output", "perf.csv") + # Update perf CSV with skipped entry update_perf_csv(exception_result=temp_file, perf_csv=output_csv) - + # Clean up temp file import os + os.unlink(temp_file) - + except Exception as e: - self.rich_console.print(f"[dim] Warning: Could not write SKIPPED status to CSV: {e}[/dim]") + self.rich_console.print( + f"[dim] Warning: Could not write SKIPPED status to CSV: {e}[/dim]" + ) def _infer_deployment_target(self, config: Dict) -> str: """ Infer deployment target from configuration structure. - + Convention over Configuration: - Presence of "k8s" or "kubernetes" field → k8s deployment - Presence of "slurm" field → slurm deployment + - ...with slurm.scheduler == "spur" (or deploy == "spur") → spur deployment - Neither present → local execution - + Args: config: Configuration dictionary - + Returns: - Deployment target: "k8s", "slurm", or "local" + Deployment target: "k8s", "spur", "slurm", or "local" """ if "k8s" in config or "kubernetes" in config: return "k8s" elif "slurm" in config: + # Spur reuses the "slurm" config block (it ships SLURM-compatible + # CLI shims), so the flavor is distinguished by an inner key rather + # than by a top-level one. + slurm_config = config.get("slurm") or {} + # An explicit deploy == "spur" is accepted too: ConfigLoader validates + # that form (build honours it), so the run path must agree or the same + # config would build for spur and run the stock SLURM template. + if ( + str(slurm_config.get("scheduler", "")).lower() == "spur" + or str(config.get("deploy", "")).lower() == "spur" + ): + return "spur" return "slurm" else: return "local" - - diff --git a/tests/unit/test_spur.py b/tests/unit/test_spur.py new file mode 100644 index 00000000..19b2e5af --- /dev/null +++ b/tests/unit/test_spur.py @@ -0,0 +1,801 @@ +#!/usr/bin/env python3 +""" +Unit tests for the spur (Crusoe) deployment backend. + +Spur ships SLURM-compatible CLI shims but cannot fan out with `srun`, so +madengine drives multi-node runs with a job ARRAY of single-node tasks that +self-form the cluster through a shared-filesystem rendezvous. These tests lock +in the contract points that make that work: + +1. `slurm.scheduler == "spur"` selects the spur backend (target inference and + ConfigLoader), and bad combinations raise. +2. `_expand_nodelist` handles both the spur (expanded) and stock SLURM + (compressed) nodelist forms. +3. The rendered job script and the slurm_multi wrapper emit array directives, + `%A_%a` log names, and the fail-fast rendezvous. +4. `SpurDeployment.monitor()` reports completion from the per-rank markers and + cannot hang when squeue is empty or unavailable. + +Copyright (c) Advanced Micro Devices, Inc. All rights reserved. +""" + +import json +import shlex +import shutil +import subprocess +from fnmatch import fnmatch +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +from madengine.deployment.base import DeploymentConfig, DeploymentStatus +from madengine.deployment.config_loader import ConfigLoader +from madengine.deployment.factory import DeploymentFactory +from madengine.deployment.slurm import SlurmDeployment +from madengine.deployment.spur import ( + DEFAULT_RENDEZVOUS_TIMEOUT, + SpurDeployment, + render_rendezvous_block, +) +from madengine.orchestration.run_orchestrator import RunOrchestrator + +# --------------------------------------------------------------------------- +# 1. Target inference + + +class TestSpurTargetInference: + """`slurm.scheduler` distinguishes spur from stock SLURM.""" + + @pytest.fixture + def orchestrator(self) -> RunOrchestrator: + args = MagicMock() + args.additional_context = None + args.live_output = True + return RunOrchestrator(args) + + @pytest.mark.parametrize( + "config,expected", + [ + ({}, "local"), + ({"slurm": {"nodes": 2}}, "slurm"), + ({"slurm": {"nodes": 2, "scheduler": "slurm"}}, "slurm"), + ({"slurm": {"nodes": 2, "scheduler": "spur"}}, "spur"), + ({"slurm": {"nodes": 2, "scheduler": "SPUR"}}, "spur"), + ({"k8s": {"namespace": "default"}}, "k8s"), + ], + ) + def test_infer_deployment_target(self, orchestrator, config, expected): + assert orchestrator._infer_deployment_target(config) == expected + + def test_runtime_context_still_wins_over_manifest(self, orchestrator): + """Regression: the manifest's build-time target must not override the + runtime --additional-context (Convention over Configuration).""" + assert orchestrator._infer_deployment_target({"k8s": {}}) == "k8s" + + @pytest.mark.parametrize( + "config,expected", + [ + ({}, "local"), + ({"slurm": {"nodes": 2}}, "slurm"), + ({"slurm": {"nodes": 2, "scheduler": "spur"}}, "spur"), + ({"deploy": "spur", "slurm": {"nodes": 2}}, "spur"), + ({"k8s": {}}, "k8s"), + ], + ) + def test_config_loader_infers_spur(self, config, expected): + assert ConfigLoader.infer_and_validate_deploy_type(config) == expected + + def test_deploy_spur_without_slurm_config_raises(self): + with pytest.raises(ValueError, match="no 'slurm' config"): + ConfigLoader.infer_and_validate_deploy_type({"deploy": "spur"}) + + def test_deploy_spur_conflicting_scheduler_raises(self): + with pytest.raises(ValueError, match="slurm.scheduler"): + ConfigLoader.infer_and_validate_deploy_type( + {"deploy": "spur", "slurm": {"scheduler": "slurm"}} + ) + + def test_unknown_scheduler_raises(self): + with pytest.raises(ValueError, match="Unknown slurm.scheduler"): + ConfigLoader.infer_and_validate_deploy_type({"slurm": {"scheduler": "pbs"}}) + + def test_spur_reuses_slurm_presets(self): + """load_config must apply the SLURM presets for spur, not fall through to local.""" + merged = ConfigLoader.load_config({"slurm": {"nodes": 2, "scheduler": "spur"}}) + assert merged["slurm"]["scheduler"] == "spur" + assert merged["slurm"]["nodes"] == 2 + # Presets contribute keys the user did not supply. + assert len(merged["slurm"]) > 2 + + +# --------------------------------------------------------------------------- +# 2. Nodelist expansion + + +class TestExpandNodelist: + """`_expand_nodelist` must work on both spur and stock SLURM forms.""" + + def test_empty(self): + assert SlurmDeployment._expand_nodelist("") == [] + + def test_spur_expanded_comma_list(self): + """Spur exposes an already-expanded list and has no `scontrol show hostnames`.""" + assert SlurmDeployment._expand_nodelist("nodeA,nodeB,nodeC") == [ + "nodeA", + "nodeB", + "nodeC", + ] + + def test_single_host(self): + assert SlurmDeployment._expand_nodelist("nodeA") == ["nodeA"] + + def test_compressed_form_uses_scontrol(self): + with patch("madengine.deployment.slurm.subprocess.run") as mock_run: + mock_run.return_value = MagicMock(returncode=0, stdout="node01\nnode02\n") + assert SlurmDeployment._expand_nodelist("node[01-02]") == [ + "node01", + "node02", + ] + + def test_compressed_form_falls_back_when_scontrol_missing(self): + with patch( + "madengine.deployment.slurm.subprocess.run", side_effect=FileNotFoundError + ): + # No expansion possible; must not raise. + assert SlurmDeployment._expand_nodelist("node[01-02]") == ["node[01-02]"] + + +# --------------------------------------------------------------------------- +# 3. Rendezvous block + + +class TestRendezvousBlock: + """The rendezvous must fail fast instead of continuing with an empty MASTER_ADDR.""" + + def test_rank0_publishes_and_peers_wait(self): + block = "\n".join(render_rendezvous_block("/shared/rendezvous", 600)) + assert "/shared/rendezvous/${SLURM_ARRAY_JOB_ID:-${SLURM_JOB_ID}}" in block + assert 'echo "$_MAD_MY_IP" > "$_MAD_REND_DIR/master_addr"' in block + assert "_MAD_REND_TIMEOUT=600" in block + + def test_timeout_writes_failure_marker_and_exits(self): + """A peer that never sees master_addr must report a failure, not hang or + continue silently: monitor() keys off the done_rank markers.""" + block = "\n".join(render_rendezvous_block("/shared/rendezvous", 600)) + assert 'if [ -z "$MASTER_ADDR" ]; then' in block + assert 'echo "1" > "$_MAD_REND_DIR/done_rank${SLURM_PROCID}"' in block + assert "exit 1" in block + assert ">&2" in block + + def test_timeout_is_configurable(self): + assert "_MAD_REND_TIMEOUT=42" in "\n".join( + render_rendezvous_block("/shared/rendezvous", 42) + ) + + +# --------------------------------------------------------------------------- +# Shared fixtures for deployment-level tests + +SPUR_MODEL_ENTRY = { + "name": "dummy_spur_model", + "url": "", + "dockerfile": "docker/dummy", + "scripts": "scripts/dummy/run.sh", + "n_gpus": "8", + "owner": "mad.support@amd.com", + "tags": ["dummy", "spur"], + "timeout": -1, + "args": "", + "env_vars": {"DOCKER_IMAGE_NAME": "registry.example.com/rocm/dummy:latest"}, +} + + +def _make_manifest(tmp_path: Path, distributed: dict, image: str = None) -> Path: + script_abs = tmp_path / SPUR_MODEL_ENTRY["scripts"] + script_abs.parent.mkdir(parents=True, exist_ok=True) + script_abs.write_text("#!/bin/bash\n# placeholder model script for unit test\n") + + image_key = image or SPUR_MODEL_ENTRY["env_vars"]["DOCKER_IMAGE_NAME"] + model_entry = dict( + SPUR_MODEL_ENTRY, + distributed=distributed, + env_vars={"DOCKER_IMAGE_NAME": image_key}, + ) + manifest = { + "built_images": { + image_key: { + "image_name": image_key, + "docker_image": image_key, + "registry_image": image_key, + }, + }, + "built_models": {image_key: model_entry}, + "context": { + "docker_env_vars": {}, + "docker_mounts": {}, + "docker_build_arg": {}, + "gpu_vendor": "AMD", + "guest_os": "UBUNTU", + "docker_gpus": "all", + }, + } + manifest_path = tmp_path / "build_manifest.json" + manifest_path.write_text(json.dumps(manifest)) + return manifest_path + + +def _make_deployment( + tmp_path: Path, cls, launcher: str, nodes: int = 3, image: str = None, **slurm_extra +): + distributed = { + "launcher": launcher, + "nnodes": nodes, + "nproc_per_node": 8, + "backend": "nccl", + "port": 29500, + } + manifest_path = _make_manifest(tmp_path, distributed, image=image) + slurm = dict( + partition="amd-rccl", + nodes=nodes, + gpus_per_node=8, + time="01:00:00", + output_dir=str(tmp_path / "slurm_results"), + exclusive=True, + **slurm_extra, + ) + additional_context = { + "gpu_vendor": "AMD", + "guest_os": "UBUNTU", + "slurm": slurm, + "distributed": distributed, + } + cfg = DeploymentConfig( + target=cls.DEPLOYMENT_TYPE, + manifest_file=str(manifest_path), + additional_context=additional_context, + ) + return cls(cfg) + + +def _render_job_script(deployment) -> str: + model_info = next(iter(deployment.manifest["built_models"].values())) + context = deployment._prepare_template_context(model_info) + return deployment.jinja_env.get_template("job.sh.j2").render(**context) + + +# --------------------------------------------------------------------------- +# 4. Rendered sbatch template + + +class TestSpurJobTemplate: + """job.sh.j2 must switch to job-array fan-out under scheduler == 'spur'.""" + + def test_spur_uses_job_array(self, tmp_path): + script = _render_job_script( + _make_deployment(tmp_path, SpurDeployment, "torchrun") + ) + assert "#SBATCH --array=0-2" in script + assert "#SBATCH --nodes=1" in script + # srun cannot fan out on spur; each array task runs its own task script. + assert 'srun bash "$TASK_SCRIPT"' not in script + assert 'bash "$TASK_SCRIPT"' in script + + def test_stock_slurm_still_uses_srun(self, tmp_path): + script = _render_job_script( + _make_deployment(tmp_path, SlurmDeployment, "torchrun") + ) + assert "#SBATCH --array" not in script + assert "#SBATCH --nodes=3" in script + + def test_spur_log_names_use_array_job_id(self, tmp_path): + """%j is each array task's own job id; result collection globs on the + array job id that sbatch returns, which is %A.""" + script = _render_job_script( + _make_deployment(tmp_path, SpurDeployment, "torchrun") + ) + assert "_%A_%a.out" in script + assert "_%A_%a.err" in script + assert "_%j_%t.out" not in script + + def test_stock_slurm_log_names_unchanged(self, tmp_path): + script = _render_job_script( + _make_deployment(tmp_path, SlurmDeployment, "torchrun") + ) + assert "_%j_%t.out" in script + assert "_%A_%a.out" not in script + + def test_spur_exports_node_rank(self, tmp_path): + """Each array task is its own single-node allocation, so SLURM_NODEID + would otherwise be 0 on every node.""" + script = _render_job_script( + _make_deployment(tmp_path, SpurDeployment, "torchrun") + ) + assert 'export SLURM_NODEID="${SLURM_ARRAY_TASK_ID:-0}"' in script + assert 'export SLURM_PROCID="${SLURM_ARRAY_TASK_ID:-0}"' in script + + def test_nodelist_helper_is_exported_to_task_script(self, tmp_path): + """The generated TASK_SCRIPT runs as a separate bash process, which does + not inherit shell functions without `export -f`.""" + script = _render_job_script( + _make_deployment(tmp_path, SlurmDeployment, "torchrun") + ) + assert "export -f mad_expand_nodelist" in script + + def test_spur_rendezvous_is_rendered(self, tmp_path): + script = _render_job_script( + _make_deployment(tmp_path, SpurDeployment, "torchrun") + ) + assert "[spur-rendezvous]" in script + assert f"_MAD_REND_TIMEOUT={DEFAULT_RENDEZVOUS_TIMEOUT}" in script + + def test_rendezvous_timeout_is_configurable(self, tmp_path): + deployment = _make_deployment( + tmp_path, SpurDeployment, "torchrun", rendezvous_timeout=1234 + ) + assert "_MAD_REND_TIMEOUT=1234" in _render_job_script(deployment) + + +# --------------------------------------------------------------------------- +# 5. slurm_multi wrapper + + +class TestSpurSlurmMultiScript: + """The slurm_multi wrapper needs the same array/rendezvous treatment.""" + + @pytest.fixture + def wrapper(self, tmp_path) -> str: + deployment = _make_deployment(tmp_path, SpurDeployment, "slurm_multi") + assert deployment.prepare() is True + return Path(deployment.script_path).read_text() + + def test_emits_array_directives(self, wrapper): + assert "#SBATCH --array=0-2" in wrapper + assert "#SBATCH --nodes=1" in wrapper + + def test_log_names_use_array_job_id(self, wrapper): + assert "_%A_%a.out" in wrapper + assert "_%j_%t.out" not in wrapper + + def test_stock_slurm_wrapper_keeps_task_log_names(self, tmp_path): + deployment = _make_deployment(tmp_path, SlurmDeployment, "slurm_multi") + assert deployment.prepare() is True + wrapper = Path(deployment.script_path).read_text() + assert "_%j_%t.out" in wrapper + assert "#SBATCH --array" not in wrapper + + def test_pulls_locally_through_a_quoted_variable(self, wrapper): + """One array task per node, so no srun fan-out for the pull.""" + assert "MAD_PULL_IMAGE=registry.example.com/rocm/dummy:latest" in wrapper + assert 'docker pull "$MAD_PULL_IMAGE"' in wrapper + assert "srun --nodes=$SLURM_NNODES" not in wrapper + + def test_docker_image_is_shell_quoted(self, tmp_path): + """Registry image names are interpolated into generated bash.""" + hostile = "registry.example.com/rocm/dummy:latest; touch /tmp/pwned" + deployment = _make_deployment( + tmp_path, SpurDeployment, "slurm_multi", image=hostile + ) + assert deployment.prepare() is True + wrapper = Path(deployment.script_path).read_text() + assert f"MAD_PULL_IMAGE={shlex.quote(hostile)}" in wrapper + assert "; touch /tmp/pwned" not in wrapper.replace(shlex.quote(hostile), "") + + def test_completion_marker_is_per_rank(self, wrapper): + """SLURM_JOB_ID is pinned to the shared array id, so the marker path + needs the rank or all N tasks race on one file.""" + assert "_rank${SLURM_ARRAY_TASK_ID:-0}.complete" in wrapper + + def test_writes_per_rank_done_marker(self, wrapper): + assert "done_rank${NODE_RANK}" in wrapper + + def test_uses_the_shared_rendezvous_block(self, wrapper): + assert "[spur-rendezvous]" in wrapper + assert 'echo "1" > "$_MAD_REND_DIR/done_rank${SLURM_PROCID}"' in wrapper + + +# --------------------------------------------------------------------------- +# 6. monitor() + + +class TestSpurMonitor: + """Marker-based completion detection, and no way to poll forever.""" + + @pytest.fixture + def deployment(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + dep._show_log_summary = MagicMock() + dep._stream_job_output = MagicMock() + return dep + + @staticmethod + def _write_markers(deployment, job_id: str, codes: dict): + marker_dir = Path(deployment.rendezvous_dir) / job_id + marker_dir.mkdir(parents=True, exist_ok=True) + for rank, code in codes.items(): + (marker_dir / f"done_rank{rank}").write_text(str(code)) + + def test_all_ranks_succeeded(self, deployment): + self._write_markers(deployment, "111", {0: 0, 1: 0, 2: 0}) + result = deployment.monitor("111") + assert result.status == DeploymentStatus.SUCCESS + deployment._show_log_summary.assert_called_once_with("111", success=True) + + def test_one_rank_failed(self, deployment): + self._write_markers(deployment, "111", {0: 0, 1: 7, 2: 0}) + result = deployment.monitor("111") + assert result.status == DeploymentStatus.FAILED + assert "1: 7" in result.message + deployment._show_log_summary.assert_called_once_with("111", success=False) + + def test_unreadable_marker_counts_as_failure(self, deployment): + self._write_markers(deployment, "111", {0: 0, 1: "garbage", 2: 0}) + assert deployment.monitor("111").status == DeploymentStatus.FAILED + + def test_partial_with_live_tasks_keeps_running(self, deployment): + self._write_markers(deployment, "111", {0: 0}) + deployment._live_task_count = MagicMock(return_value=2) + result = deployment.monitor("111") + assert result.status == DeploymentStatus.RUNNING + assert "1/3 ranks done" in result.message + + def test_startup_grace_before_tasks_are_registered(self, deployment): + """Right after sbatch, spur's squeue lists nothing; that must not be + mistaken for a dead array.""" + deployment._live_task_count = MagicMock(return_value=0) + for _ in range(deployment._SPUR_DEAD_POLLS + 2): + assert deployment.monitor("111").status == DeploymentStatus.RUNNING + + def test_startup_grace_is_bounded(self, deployment): + """An array that dies before squeue ever lists it writes no marker and + never appears: the startup window must still end, or monitor() (which the + caller polls without a timeout) never returns.""" + deployment._live_task_count = MagicMock(return_value=0) + for _ in range(deployment._SPUR_STARTUP_POLLS - 1): + assert deployment.monitor("111").status == DeploymentStatus.RUNNING + result = deployment.monitor("111") + assert result.status == DeploymentStatus.FAILED + assert "was ever seen in the queue" in result.message + deployment._show_log_summary.assert_called_once_with("111", success=False) + + def test_startup_counter_resets_once_tasks_appear(self, deployment): + """A slow queue must not accumulate toward the startup bound.""" + deployment._live_task_count = MagicMock( + side_effect=[0] * (deployment._SPUR_STARTUP_POLLS - 1) + + [2] + + [0] * (deployment._SPUR_STARTUP_POLLS + 2) + ) + for _ in range(deployment._SPUR_STARTUP_POLLS): + assert deployment.monitor("111").status == DeploymentStatus.RUNNING + # Seen alive, then empty again: now the (shorter) dead-array window applies. + for _ in range(deployment._SPUR_DEAD_POLLS - 1): + assert deployment.monitor("111").status == DeploymentStatus.RUNNING + assert deployment.monitor("111").status == DeploymentStatus.FAILED + + def test_dead_array_fails_after_grace_window(self, deployment): + self._write_markers(deployment, "111", {0: 0}) + deployment._live_task_count = MagicMock(side_effect=[3] + [0] * 10) + assert deployment.monitor("111").status == DeploymentStatus.RUNNING # seen live + for _ in range(deployment._SPUR_DEAD_POLLS - 1): + assert deployment.monitor("111").status == DeploymentStatus.RUNNING + result = deployment.monitor("111") + assert result.status == DeploymentStatus.FAILED + assert "1/3 ranks reported completion" in result.message + + def test_persistent_squeue_outage_gives_up(self, deployment): + """live == -1 must not reset the loop forever: the caller polls without + a timeout, so an unreachable control plane would hang the run.""" + deployment._live_task_count = MagicMock(return_value=-1) + for _ in range(deployment._SPUR_UNKNOWN_POLLS - 1): + assert deployment.monitor("111").status == DeploymentStatus.RUNNING + assert deployment.monitor("111").status == DeploymentStatus.UNKNOWN + + def test_markers_win_over_squeue_outage(self, deployment): + deployment._live_task_count = MagicMock(return_value=-1) + self._write_markers(deployment, "111", {0: 0, 1: 0, 2: 0}) + assert deployment.monitor("111").status == DeploymentStatus.SUCCESS + + def test_live_output_streams_instead_of_summary(self, deployment): + deployment.config.additional_context["live_output"] = True + deployment._live_task_count = MagicMock(return_value=3) + deployment.monitor("111") + deployment._stream_job_output.assert_called_with("111") + + self._write_markers(deployment, "111", {0: 0, 1: 0, 2: 0}) + deployment.monitor("111") + deployment._stream_job_output.assert_called_with("111", final=True) + deployment._show_log_summary.assert_not_called() + + +class TestLiveTaskCount: + """squeue parsing for the liveness guard.""" + + @pytest.fixture + def deployment(self, tmp_path): + return _make_deployment(tmp_path, SpurDeployment, "torchrun") + + def _run_squeue(self, deployment, stdout, returncode=0): + with patch("madengine.deployment.spur.subprocess.run") as mock_run: + mock_run.return_value = MagicMock(returncode=returncode, stdout=stdout) + count = deployment._live_task_count("111", "madengine-dummy_spur_model") + return count, mock_run.call_args[0][0] + + def test_counts_only_my_array_tasks(self, deployment): + """A concurrent run of the same model must not inflate the count.""" + stdout = ( + "111_0 madengine-dummy_spur_model RUNNING\n" + "111_1 madengine-dummy_spur_model RUNNING\n" + "222_0 madengine-dummy_spur_model RUNNING\n" + ) + count, _ = self._run_squeue(deployment, stdout) + assert count == 2 + + def test_ignores_finished_states(self, deployment): + stdout = ( + "111_0 madengine-dummy_spur_model COMPLETED\n" + "111_1 madengine-dummy_spur_model RUNNING\n" + ) + count, _ = self._run_squeue(deployment, stdout) + assert count == 1 + + def test_pending_array_range(self, deployment): + count, _ = self._run_squeue( + deployment, "111_[0-2] madengine-dummy_spur_model PENDING\n" + ) + assert count == 1 + + def test_falls_back_to_name_when_no_id_matches(self, deployment): + count, _ = self._run_squeue( + deployment, "999 madengine-dummy_spur_model RUNNING\n" + ) + assert count == 1 + + def test_non_zero_exit_is_unknown(self, deployment): + count, _ = self._run_squeue(deployment, "", returncode=1) + assert count == -1 + + def test_exception_is_unknown(self, deployment): + with patch("madengine.deployment.spur.subprocess.run", side_effect=OSError): + assert deployment._live_task_count("111", "madengine-x") == -1 + + def test_unset_user_omits_the_flag(self, deployment): + """`squeue -u ""` is an error, so drop -u entirely.""" + with patch.dict("os.environ", {}, clear=True): + _, cmd = self._run_squeue(deployment, "") + assert "-u" not in cmd + + def test_user_is_passed_when_set(self, deployment): + with patch.dict("os.environ", {"USER": "someone"}, clear=True): + _, cmd = self._run_squeue(deployment, "") + assert cmd[cmd.index("-u") + 1] == "someone" + + +# --------------------------------------------------------------------------- +# 7. Backend selection wiring + + +class TestFactoryRegistration: + """The inferred target string must reach the right class.""" + + def test_spur_target_creates_spur_deployment(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + assert isinstance(DeploymentFactory.create(dep.config), SpurDeployment) + + def test_slurm_target_still_creates_slurm_deployment(self, tmp_path): + dep = _make_deployment(tmp_path, SlurmDeployment, "torchrun") + created = DeploymentFactory.create(dep.config) + assert isinstance(created, SlurmDeployment) + assert not isinstance(created, SpurDeployment) + + def test_stock_slurm_is_not_flagged_as_spur(self): + """IS_SPUR gates every spur branch inside the shared SLURM code.""" + assert SlurmDeployment.IS_SPUR is False + assert SpurDeployment.IS_SPUR is True + + +class TestInferenceConsistency: + """The three inference sites (ConfigLoader, build, run) must agree, or a + config builds for one backend and runs on the other.""" + + @pytest.fixture + def orchestrator(self) -> RunOrchestrator: + args = MagicMock() + args.additional_context = None + args.live_output = True + return RunOrchestrator(args) + + @pytest.mark.parametrize( + "config", + [ + {}, + {"slurm": {"nodes": 2}}, + {"slurm": {"nodes": 2, "scheduler": "slurm"}}, + {"slurm": {"nodes": 2, "scheduler": "spur"}}, + {"deploy": "spur", "slurm": {"nodes": 2}}, + {"deploy": "slurm", "slurm": {"nodes": 2}}, + {"k8s": {"namespace": "default"}}, + ], + ) + def test_config_loader_and_run_orchestrator_agree(self, orchestrator, config): + assert orchestrator._infer_deployment_target( + config + ) == ConfigLoader.infer_and_validate_deploy_type(config) + + +class TestSpurValidate: + """spur's scontrol is only partially implemented, so it must not be required.""" + + def test_validate_does_not_probe_scontrol(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + with patch("madengine.deployment.slurm.subprocess.run") as mock_run: + mock_run.return_value = MagicMock(returncode=0) + assert dep.validate() is True + probed = [call.args[0] for call in mock_run.call_args_list] + assert ["which", "scontrol"] not in probed + assert ["which", "sbatch"] in probed + + def test_stock_slurm_still_requires_scontrol(self, tmp_path): + dep = _make_deployment(tmp_path, SlurmDeployment, "torchrun") + with patch("madengine.deployment.slurm.subprocess.run") as mock_run: + mock_run.return_value = MagicMock(returncode=0) + assert dep.validate() is True + probed = [call.args[0] for call in mock_run.call_args_list] + assert ["which", "scontrol"] in probed + + +# --------------------------------------------------------------------------- +# 8. Node health preflight + + +class TestNodePreflight: + """The srun-based preflight cannot work on spur, and the nodelist it pins is + actively harmful there: every array task requests --nodes=1, so a multi-node + #SBATCH --nodelist would make each task demand all of them.""" + + @staticmethod + def _submit(deployment): + with patch("madengine.deployment.slurm.subprocess.run") as mock_run: + mock_run.return_value = MagicMock( + returncode=0, stdout="Submitted batch job 4242", stderr="" + ) + return deployment.deploy() + + def test_spur_skips_preflight_and_submits(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + assert dep.prepare() is True + with patch("madengine.deployment.slurm.SlurmNodeSelector") as selector: + result = self._submit(dep) + selector.assert_not_called() + assert result.deployment_id == "4242" + assert "nodelist" not in dep.slurm_config + assert "#SBATCH --nodelist" not in Path(dep.script_path).read_text() + + def test_stock_slurm_preflight_is_unchanged(self, tmp_path): + """Regression: SLURM still health-checks and still gates multi-node + submission when there are not enough clean nodes.""" + dep = _make_deployment(tmp_path, SlurmDeployment, "torchrun") + assert dep.prepare() is True + with patch("madengine.deployment.slurm.SlurmNodeSelector") as selector: + selector.return_value.select_nodes.return_value = (["nodeA"], "") + result = self._submit(dep) + selector.return_value.select_nodes.assert_called_once() + assert result.status == DeploymentStatus.FAILED + assert "Not enough clean nodes" in result.message + + def test_stock_slurm_preflight_still_pins_clean_nodes(self, tmp_path): + dep = _make_deployment(tmp_path, SlurmDeployment, "torchrun") + assert dep.prepare() is True + with patch("madengine.deployment.slurm.SlurmNodeSelector") as selector: + selector.return_value.select_nodes.return_value = ( + ["nodeA", "nodeB", "nodeC"], + "", + ) + result = self._submit(dep) + assert result.deployment_id == "4242" + assert dep.slurm_config["nodelist"] == "nodeA,nodeB,nodeC" + assert ( + "#SBATCH --nodelist=nodeA,nodeB,nodeC" in Path(dep.script_path).read_text() + ) + + +# --------------------------------------------------------------------------- +# 9. Generated scripts are valid bash + + +@pytest.mark.skipif(shutil.which("bash") is None, reason="bash not available") +class TestGeneratedScriptsAreValidBash: + """The jinja branches must not produce a script the shell rejects - a syntax + error would only surface as a failed job on the cluster.""" + + @pytest.mark.parametrize("cls", [SlurmDeployment, SpurDeployment]) + def test_rendered_template_parses(self, tmp_path, cls): + script = tmp_path / "job.sh" + script.write_text( + _render_job_script(_make_deployment(tmp_path, cls, "torchrun")) + ) + assert subprocess.run(["bash", "-n", str(script)]).returncode == 0 + + @pytest.mark.parametrize("cls", [SlurmDeployment, SpurDeployment]) + def test_slurm_multi_wrapper_parses(self, tmp_path, cls): + deployment = _make_deployment(tmp_path, cls, "slurm_multi") + assert deployment.prepare() is True + assert subprocess.run(["bash", "-n", deployment.script_path]).returncode == 0 + + +# --------------------------------------------------------------------------- +# 10. Log/artifact collection keys on the id sbatch returned + + +class TestLogCollectionCompatibility: + """collect_results()/_show_log_summary() glob on the deployment id, which for + an array is the array job id (%A), not each task's own job id (%j).""" + + @staticmethod + def _sbatch_log_names(script: str, job_id: str) -> list: + return [ + line.split("=", 1)[1] + .replace("%A", job_id) + .replace("%a", "0") + .replace("%j", "999") + .replace("%t", "0") + for line in script.splitlines() + if line.startswith("#SBATCH --output=") + ] + + def test_spur_sbatch_log_names_match_the_collection_glob(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + names = self._sbatch_log_names(_render_job_script(dep), "4242") + assert names + for name in names: + assert fnmatch(Path(name).name, "madengine-*_4242_*.out") + + def test_show_log_summary_finds_spur_array_logs(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + dep.output_dir.mkdir(parents=True, exist_ok=True) + log = dep.output_dir / "madengine-dummy_spur_model_4242_1.out" + log.write_text("done\n") + dep.console = MagicMock() + dep._show_log_summary("4242", success=True) + printed = " ".join(str(c.args[0]) for c in dep.console.print.call_args_list) + assert str(log) in printed + + @pytest.mark.parametrize("cls", [SlurmDeployment, SpurDeployment]) + def test_node_logs_key_on_the_collection_job_id(self, tmp_path, cls): + """The per-node logs the task script writes are the ones collect_results + reads for multi-node runs, so they must carry the same id.""" + script = _render_job_script(_make_deployment(tmp_path, cls, "torchrun")) + assert ( + 'export MAD_COLLECT_JOB_ID="${SLURM_ARRAY_JOB_ID:-$SLURM_JOB_ID}"' in script + ) + assert "_${MAD_COLLECT_JOB_ID}_node_${SLURM_PROCID}.out" in script + assert "_${SLURM_JOB_ID}_node_${SLURM_PROCID}.out" not in script + + +# --------------------------------------------------------------------------- +# 11. Rendezvous configuration plumbing + + +class TestRendezvousConfig: + def test_rendezvous_dir_is_under_the_shared_output_dir(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + assert Path(dep.rendezvous_dir) == dep.output_dir.resolve() / "spur_rendezvous" + + def test_default_timeout(self, tmp_path): + dep = _make_deployment(tmp_path, SpurDeployment, "torchrun") + assert dep.rendezvous_timeout == DEFAULT_RENDEZVOUS_TIMEOUT + + def test_timeout_comes_from_the_slurm_block(self, tmp_path): + dep = _make_deployment( + tmp_path, SpurDeployment, "torchrun", rendezvous_timeout="60" + ) + assert dep.rendezvous_timeout == 60 + + def test_template_context_carries_the_scheduler_flavor(self, tmp_path): + def model(d): + return next(iter(d.manifest["built_models"].values())) + + spur = _make_deployment(tmp_path / "spur", SpurDeployment, "torchrun") + slurm = _make_deployment(tmp_path / "slurm", SlurmDeployment, "torchrun") + assert spur._prepare_template_context(model(spur))["scheduler"] == "spur" + assert slurm._prepare_template_context(model(slurm))["scheduler"] == "slurm"