From 39d177fc9d399ea77d0d47c92b2bc06de1664c14 Mon Sep 17 00:00:00 2001 From: johnsonms Date: Wed, 10 Jun 2026 00:43:09 +0000 Subject: [PATCH 1/3] ci(fa4): assert cute dep floors in CI; fail loudly on a stale SIF MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit run_fa4_ci.py installs FA4 with --no-deps (to keep the SIF's baked torch/cudnn), so the nvidia-cutlass-dsl>=4.5.2 / quack-kernels>=0.5.0 floors in flash_attn/cute/pyproject.toml are not enforced at install time. A SIF baked before a floor bump keeps a stale dep — e.g. cutlass-dsl 4.4.2, which can't convert the AuxData JIT arg and dies with a cryptic DSLRuntimeError deep in SM100 kernel launch (reproduced on B200). Upgrading the dep in-place is not viable: the --writable-tmpfs overlay is RAM-backed and too small for a cutlass-dsl reinstall (ENOSPC, and a partial removal corrupts the baked torch). So instead of installing, add assert_dsl_floor.py — it reads the floors from pyproject (no hardcoded version to drift) and fails with an actionable "rebake the image" message when the installed cutlass-dsl/quack are below them. Wired into run_step right after the editable install. The durable fix is to rebake the image at the current floors and bump the digest in .github/workflows/ci.yml; this guard makes future drift fail fast instead of silently. --- tools/ci/assert_dsl_floor.py | 65 ++++++++++++++++++++++++++++++++++++ tools/ci/run_fa4_ci.py | 14 +++++++- 2 files changed, 78 insertions(+), 1 deletion(-) create mode 100644 tools/ci/assert_dsl_floor.py diff --git a/tools/ci/assert_dsl_floor.py b/tools/ci/assert_dsl_floor.py new file mode 100644 index 00000000000..ee0035447d4 --- /dev/null +++ b/tools/ci/assert_dsl_floor.py @@ -0,0 +1,65 @@ +#!/usr/bin/env python3 +"""Fail loudly if the CI image's deps are below the flash_attn/cute/pyproject.toml floors. + +Runs inside the SIF before tests. The FA4 install in run_fa4_ci.py uses --no-deps (to keep the +SIF's torch/cudnn), so pyproject floors are not enforced at install time. A SIF baked before a +floor bump therefore keeps a stale dep — e.g. nvidia-cutlass-dsl 4.4.2, which can't convert the +AuxData JIT arg and dies with a cryptic DSLRuntimeError deep in SM100 kernel launch. This check +turns that into an actionable "rebake the image" message up front. + +Reads the floor from pyproject so there is no hardcoded version here to drift out of sync. +""" + +from __future__ import annotations + +import sys +import tomllib +from importlib.metadata import PackageNotFoundError, version + +from packaging.requirements import Requirement +from packaging.version import Version + +# Deps whose floor a stale SIF is known to silently violate. Other pyproject deps (torch, einops…) +# are baked to match the image and not version-sensitive in the same way, so we don't gate on them. +CHECKED = ("nvidia-cutlass-dsl", "quack-kernels") + + +def main(pyproject_path: str) -> int: + with open(pyproject_path, "rb") as f: + deps = tomllib.load(f)["project"]["dependencies"] + reqs = {r.name: r for r in (Requirement(d) for d in deps) if r.name in CHECKED} + + failures: list[str] = [] + oks: list[str] = [] + for name in CHECKED: + req = reqs.get(name) + if req is None: + continue # not a hard dep in this pyproject — nothing to enforce + try: + installed = version(name) + except PackageNotFoundError: + failures.append(f"{name}: not installed (floor {req.specifier})") + continue + if req.specifier.contains(Version(installed), prereleases=True): + oks.append(f"{name}={installed}") + else: + failures.append(f"{name}: installed {installed} does not satisfy floor {req.specifier}") + + if failures: + print("ERROR: CI image deps are below the flash_attn/cute/pyproject.toml floor:", file=sys.stderr) + for line in failures: + print(f" - {line}", file=sys.stderr) + print( + "\nThe SIF was likely baked before a floor bump. Rebake the image " + "(tools/ci/docker/build.sh + tag_and_push.sh) and update the digest in " + ".github/workflows/ci.yml.", + file=sys.stderr, + ) + return 1 + + print("DSL floor check OK: " + ", ".join(oks)) + return 0 + + +if __name__ == "__main__": + sys.exit(main(sys.argv[1] if len(sys.argv) > 1 else "flash_attn/cute/pyproject.toml")) diff --git a/tools/ci/run_fa4_ci.py b/tools/ci/run_fa4_ci.py index e539df7f056..0182c505298 100644 --- a/tools/ci/run_fa4_ci.py +++ b/tools/ci/run_fa4_ci.py @@ -106,9 +106,21 @@ def run_step(step: Step, repo_root: Path, base_env: dict[str, str], sif: str, wo print(f"=== {step.name} ===") # Install FA4 from the current repo inside this exec invocation. + # --no-deps keeps the SIF's baked torch/cudnn; deps are expected to already satisfy the + # pyproject floors via the image (rebuilt with tools/ci/docker/build.sh). We do NOT upgrade + # deps here: the --writable-tmpfs overlay is RAM-backed and too small to hold a cutlass-dsl + # reinstall (it ENOSPCs and can corrupt the baked torch). Instead assert_dsl_floor.py below + # fails loudly if the image is stale, pointing at a rebake. # Must be done per-step because --writable-tmpfs creates a fresh overlay each time. install_cmd = f"uv pip install --system --break-system-packages --no-deps -q -e {shlex.quote(str(repo_root / 'flash_attn/cute'))}" + # Guard against a SIF baked with deps below the pyproject floor (the silent --no-deps gap that + # otherwise surfaces as a cryptic DSLRuntimeError on the SM100 path). Cheap: reads versions, no install. + floor_check_cmd = ( + f"python3 {shlex.quote(str(repo_root / 'tools/ci/assert_dsl_floor.py'))} " + f"{shlex.quote(str(repo_root / 'flash_attn/cute/pyproject.toml'))}" + ) + # Convert relative test/benchmark paths to absolute so we can run from /tmp. # Running from /tmp ensures Python does not insert repo_root into sys.path[0] # (which would cause flash_attn/__init__.py to trigger FA2 imports unavailable in the SIF). @@ -118,7 +130,7 @@ def run_step(step: Step, repo_root: Path, base_env: dict[str, str], sif: str, wo ] env_exports = " && ".join(f"export {k}={shlex.quote(v)}" for k, v in step.extra_env.items()) inner_cmd = shlex.join(command) - shell_parts = [install_cmd] + shell_parts = [install_cmd, floor_check_cmd] if env_exports: shell_parts.append(env_exports) shell_parts.append(f"cd /tmp && {inner_cmd}") From bdaeaed0fe903703142f32afa5ad979a64a1b87b Mon Sep 17 00:00:00 2001 From: Johnsonms Date: Wed, 10 Jun 2026 01:12:40 +0000 Subject: [PATCH 2/3] ci(fa4): bump cu130 image to 26.06.10 (cutlass-dsl 4.5.2 / quack 0.5.0) --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f552c83bb8a..5c718d78683 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -45,4 +45,4 @@ jobs: with: test-filter: ${{ env.FA4_TEST_FILTER }} fa4_image_cu129: "togethercomputer/training-performance:flash-attn-cu12.9-26.03.25@sha256:304a5c3d2b3a75b151cd2a964cd26d444e0d8b5686d63943df13378c9705f943" - fa4_image_cu130: "togethercomputer/training-performance:flash-attn-cu13.0-26.04.01@sha256:56e50b056eb4d671410846c3483e843ee7bd0f5b13cb45b6f0d7eb8bd27694a5" + fa4_image_cu130: "togethercomputer/training-performance:flash-attn-cu13.0-26.06.10@sha256:f1efd03b9d78cf65d9f8df107d2f6f6d0a464cb8b773fd9364765f22f4772006" From 00e8933b51c0294669c33c690e05508d50a3a2dd Mon Sep 17 00:00:00 2001 From: Johnsonms Date: Wed, 10 Jun 2026 01:59:17 +0000 Subject: [PATCH 3/3] ci(fa4): fall back to tomli when tomllib is unavailable (Python 3.10) --- tools/ci/assert_dsl_floor.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/tools/ci/assert_dsl_floor.py b/tools/ci/assert_dsl_floor.py index ee0035447d4..19a2f2f4312 100644 --- a/tools/ci/assert_dsl_floor.py +++ b/tools/ci/assert_dsl_floor.py @@ -13,12 +13,22 @@ from __future__ import annotations import sys -import tomllib from importlib.metadata import PackageNotFoundError, version from packaging.requirements import Requirement from packaging.version import Version +try: + import tomllib # Python 3.11+ +except ModuleNotFoundError: # Python 3.10 (pyproject declares requires-python >=3.10) + try: + import tomli as tomllib + except ModuleNotFoundError: + sys.exit( + "ERROR: assert_dsl_floor.py needs a TOML parser — use Python 3.11+ (stdlib tomllib) " + "or `pip install tomli` on 3.10." + ) + # Deps whose floor a stale SIF is known to silently violate. Other pyproject deps (torch, einops…) # are baked to match the image and not version-sensitive in the same way, so we don't gate on them. CHECKED = ("nvidia-cutlass-dsl", "quack-kernels")