From e0a1f050591276cf8a1b5ee039fc272b06a32ea6 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 07:15:17 +0000 Subject: [PATCH 01/14] feat: (jax) Add Build JAX wheels in Multi-Arch CI --- .../multi_arch_build_linux_jax_wheels_ci.yml | 161 +++++++++++++++--- .github/workflows/multi_arch_ci.yml | 6 +- .github/workflows/multi_arch_ci_linux.yml | 34 +++- .../configure_jax_release_matrix.py | 45 ++++- .../github_actions/configure_multi_arch_ci.py | 13 +- 5 files changed, 225 insertions(+), 34 deletions(-) diff --git a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml index 218311d66ec..141fea57724 100644 --- a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml +++ b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml @@ -1,11 +1,15 @@ # Copyright Advanced Micro Devices, Inc. # SPDX-License-Identifier: MIT -# Placeholder reusable workflow for the Linux JAX wheel CI flow. +# Reusable single-build workflow for the Linux JAX wheel CI flow. # -# This file intentionally does not build or upload JAX wheels yet. -# It exists to land the reusable workflow interface first, so follow-up PRs can -# wire and test the real JAX CI build implementation incrementally. +# This workflow is called once per CI matrix cell from multi_arch_ci_linux.yml. +# +# Built shape: +# - One manylinux JAX wheel build for one python_version and jax_ref matrix cell. +# - Installs ROCm packages from the run-scoped CI find-links URL. +# - Uploads the resulting wheels to run-scoped CI Python artifacts. +# - Build-only for the first CI rollout; test wiring should be added separately. name: Multi-Arch Build Linux JAX Wheels (CI) @@ -64,7 +68,6 @@ on: value: ${{ jobs.build_jax_wheels.outputs.jax_plugin_version }} jax_pjrt_version: value: ${{ jobs.build_jax_wheels.outputs.jax_pjrt_version }} - workflow_dispatch: inputs: artifact_group: @@ -109,40 +112,142 @@ on: description: "Branch, tag, or SHA to checkout. Defaults to the triggering ref." type: string default: "" - -run-name: Placeholder Multi-Arch Linux JAX Wheels CI (${{ inputs.artifact_group }}, ${{ inputs.rocm_version }}, py${{ inputs.python_version }}, ${{ inputs.jax_ref }}) - +run-name: Build Multi-Arch Linux JAX Wheels CI (${{ inputs.artifact_group }}, ${{ inputs.rocm_version }}, py${{ inputs.python_version }}, ${{ inputs.jax_ref }}) permissions: contents: read + id-token: write jobs: build_jax_wheels: - name: Placeholder | py ${{ inputs.python_version }} | jax ${{ inputs.jax_ref }} - runs-on: ubuntu-24.04 + name: Build | py ${{ inputs.python_version }} | jax ${{ inputs.jax_ref }} + runs-on: ${{ github.repository_owner == 'ROCm' && 'azure-linux-scale-rocm' || 'ubuntu-24.04' }} outputs: - package_find_links_url: ${{ steps.placeholder_outputs.outputs.package_find_links_url }} - jax_version: ${{ steps.placeholder_outputs.outputs.jax_version }} - jaxlib_version: ${{ steps.placeholder_outputs.outputs.jaxlib_version }} - jax_plugin_version: ${{ steps.placeholder_outputs.outputs.jax_plugin_version }} - jax_pjrt_version: ${{ steps.placeholder_outputs.outputs.jax_pjrt_version }} + package_find_links_url: ${{ steps.upload_jax_wheels.outputs.package_find_links_url }} + jax_version: ${{ steps.write_jax_versions.outputs.jax_version }} + jaxlib_version: ${{ steps.write_jax_versions.outputs.jaxlib_version }} + jax_plugin_version: ${{ steps.write_jax_versions.outputs.jax_plugin_version }} + jax_pjrt_version: ${{ steps.write_jax_versions.outputs.jax_pjrt_version }} + env: + MANYLINUX_IMAGE_TAG: therock-jax-manylinux:${{ inputs.jax_ref }}-${{ inputs.python_version }} + RELEASE_TYPE: ci steps: - - name: Validate placeholder inputs + - name: Checkout TheRock + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + repository: ${{ inputs.repository || github.repository }} + ref: ${{ inputs.ref || github.ref_name }} + + - name: Checkout rocm-jax + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + path: rocm-jax + repository: ROCm/rocm-jax + ref: ${{ inputs.jax_ref }} + + - name: Checkout JAX + uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 + with: + path: jax-source + repository: ${{ inputs.jax_repository }} + ref: ${{ inputs.jax_ref }} + + + - name: Configure Git Identity + run: | + git config --global user.name "therockbot" + git config --global user.email "therockbot@amd.com" + + - name: Set up Python + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6.3.0 + with: + python-version: ${{ inputs.python_version }} + + - name: Install python deps for CI + run: | + pip install -r external-builds/jax/requirements-jax.txt + + - name: Validate CI build mode run: | if [ "${{ inputs.build_mode }}" != "manylinux" ]; then - echo "Only build_mode=manylinux is supported by the planned JAX CI workflow." + echo "Only build_mode=manylinux is supported by this CI workflow." exit 1 fi - - name: Emit placeholder outputs - id: placeholder_outputs + - name: Set package dist dir + run: | + echo "PACKAGE_DIST_DIR=${GITHUB_WORKSPACE}/jax-source/dist" >> "$GITHUB_ENV" + + - name: Build manylinux image + working-directory: jax-source + env: + THEROCK_VERSION: "" + GFX_ARCH: ${{ inputs.gfx_arch }} + ROCM_PACKAGE_FIND_LINKS_URL: ${{ inputs.rocm_package_find_links_url }} + run: | + cp ../rocm-jax/docker/manylinux/Dockerfile.jax-manylinux_2_28-therock \ + Dockerfile.jax-manylinux_2_28-therock + + docker build \ + -t "${MANYLINUX_IMAGE_TAG}" \ + --file=Dockerfile.jax-manylinux_2_28-therock \ + --build-arg=THEROCK_INDEX_URL="${ROCM_PACKAGE_FIND_LINKS_URL}" \ + --build-arg=THEROCK_VERSION="${THEROCK_VERSION}" \ + --build-arg=GFX_ARCH="${GFX_ARCH}" \ + --progress=plain \ + . + + - name: Determine wheel version suffix + run: | + python build_tools/github_actions/determine_version.py \ + --rocm-version "${{ inputs.rocm_version }}" + + - name: Build JAX Wheels in manylinux container + working-directory: jax-source + env: + ROCM_VERSION: ${{ inputs.rocm_version }} + PYTHON_VERSION: ${{ inputs.python_version }} + run: | + docker run --rm \ + --user root \ + --env ROCM_VERSION="${ROCM_VERSION}" \ + --env PYTHON_VERSION="${PYTHON_VERSION}" \ + --env ML_WHEEL_VERSION_SUFFIX="${version_suffix}" \ + --env PACKAGE_DIST_DIR="/workspace/jax-source/dist" \ + --volume "${GITHUB_WORKSPACE}:/workspace" \ + --workdir /workspace/jax-source \ + "${MANYLINUX_IMAGE_TAG}" \ + bash -lc ' + python build/build.py build \ + --wheels=jax-rocm-plugin,jax-rocm-pjrt \ + --python_version="${PYTHON_VERSION}" \ + --bazel_startup_options=--bazelrc=build/rocm/rocm.bazelrc \ + --bazel_options=--config=rocm_release_wheel \ + --bazel_options=--repo_env=ROCM_PATH=$(rocm-sdk path --root) \ + --bazel_options=--repo_env=ML_WHEEL_TYPE=release \ + --bazel_options=--repo_env=ML_WHEEL_VERSION_SUFFIX="${ML_WHEEL_VERSION_SUFFIX}" \ + --bazel_options=--//jaxlib/tools:jaxlib_git_hash=$(git rev-parse HEAD) \ + --verbose \ + --detailed_timestamped_log \ + --output_path=$(pwd)/dist + ' + + - name: Extract JAX versions from built wheels + id: write_jax_versions + run: | + python3 ./build_tools/github_actions/write_jax_versions.py \ + --dist-dir "${{ env.PACKAGE_DIST_DIR }}" + + - name: Configure AWS Credentials + uses: ./.github/actions/configure_aws_artifacts_credentials + with: + release_type: ci + + - name: Upload JAX wheels to CI artifacts + id: upload_jax_wheels run: | - echo "::notice::Placeholder workflow only. No JAX wheels are built or uploaded." - - { - echo "package_find_links_url=" - echo "jax_version=" - echo "jaxlib_version=" - echo "jax_plugin_version=" - echo "jax_pjrt_version=" - } >> "${GITHUB_OUTPUT}" + python build_tools/github_actions/upload_python_packages.py \ + --input-packages-dir="${{ env.PACKAGE_DIST_DIR }}" \ + --artifact-group="${{ inputs.artifact_group }}" \ + --run-id="${{ github.run_id }}" \ + --multiarch diff --git a/.github/workflows/multi_arch_ci.yml b/.github/workflows/multi_arch_ci.yml index 0f284772705..a26f8663ce8 100644 --- a/.github/workflows/multi_arch_ci.yml +++ b/.github/workflows/multi_arch_ci.yml @@ -50,6 +50,10 @@ on: type: boolean default: true description: "Build PyTorch wheels" + build_jax: + type: boolean + default: true + description: "Build JAX wheels" pull_request: types: - labeled @@ -83,7 +87,7 @@ jobs: prebuilt_stages: ${{ inputs.prebuilt_stages || '' }} baseline_run_id: ${{ inputs.baseline_run_id || '' }} build_pytorch: ${{ github.event_name != 'workflow_dispatch' || inputs.build_pytorch }} - build_jax: false + build_jax: ${{ github.event_name != 'workflow_dispatch' || inputs.build_jax }} linux_build_and_test: name: Linux::${{ fromJSON(needs.setup.outputs.linux_build_config || '{}').build_variant_label || 'skip' }} diff --git a/.github/workflows/multi_arch_ci_linux.yml b/.github/workflows/multi_arch_ci_linux.yml index f9ff5ba0bc4..34ce4fc44db 100644 --- a/.github/workflows/multi_arch_ci_linux.yml +++ b/.github/workflows/multi_arch_ci_linux.yml @@ -12,9 +12,9 @@ on: JSON object with build configuration for this platform. Fields: artifact_group, per_family_info, dist_amdgpu_families, build_variant_label, build_variant_cmake_preset, - build_variant_suffix, build_pytorch, + build_variant_suffix, build_pytorch, build_jax, build_native_linux, test_python_packages_matrix, - pytorch_build_matrix, + pytorch_build_matrix, jax_build_matrix, prebuilt_stages, baseline_run_id. test_labels: type: string @@ -320,3 +320,33 @@ jobs: permissions: contents: read id-token: write + + build_jax_wheels: + needs: [build_python_packages] + name: Build JAX Wheels + if: >- + ${{ + !failure() && + !cancelled() && + fromJSON(inputs.build_config).build_jax == true && + toJSON(fromJSON(inputs.build_config).jax_build_matrix) != '[]' + }} + strategy: + fail-fast: false + matrix: + include: ${{ fromJSON(inputs.build_config).jax_build_matrix }} + uses: ./.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml + with: + artifact_group: ${{ fromJSON(inputs.build_config).artifact_group }} + python_version: ${{ matrix.python_version }} + jax_repository: ${{ matrix.jax_repository }} + build_mode: ${{ matrix.build_mode }} + gfx_arch: ${{ matrix.gfx_arch }} + jax_ref: ${{ matrix.jax_ref }} + rocm_version: ${{ inputs.rocm_package_version }} + rocm_package_find_links_url: ${{ needs.build_python_packages.outputs.package_find_links_url }} + repository: ${{ inputs.repository }} + ref: ${{ inputs.ref }} + permissions: + contents: read + id-token: write diff --git a/build_tools/github_actions/configure_jax_release_matrix.py b/build_tools/github_actions/configure_jax_release_matrix.py index e6e034a8bd3..2633c61a83c 100644 --- a/build_tools/github_actions/configure_jax_release_matrix.py +++ b/build_tools/github_actions/configure_jax_release_matrix.py @@ -2,7 +2,7 @@ # Copyright Advanced Micro Devices, Inc. # SPDX-License-Identifier: MIT -"""Generate JAX release build matrix for workflows.""" +"""Generate JAX build matrices for CI and release workflows.""" import argparse import json @@ -36,6 +36,9 @@ }, ] +CI_JAX_REFS = {"rocm-jaxlib-v0.10.0", "rocm-jaxlib-v0.10.2"} +CI_PYTHON_VERSIONS = PYTHON_VERSIONS + def _split_values(raw: str) -> list[str]: """Split comma, semicolon, or whitespace-separated workflow input values.""" @@ -65,6 +68,46 @@ def generate_jax_matrix( return matrix +def generate_jax_matrix_for_ci( + python_versions: list[str] | None, +) -> list[dict[str, str]]: + """Generate the JAX matrix used by Multi-Arch CI. + + The release matrix includes legacy/native JAX entries such as + rocm-jaxlib-v0.9.1. The CI workflow is manylinux-only, so filter the + release matrix to the supported manylinux entries. + """ + versions = python_versions if python_versions else CI_PYTHON_VERSIONS + + matrix: list[dict[str, str]] = [] + for cell in generate_jax_matrix(versions): + build_mode = str(cell["build_mode"]) + jax_ref = str(cell["jax_ref"]) + + if build_mode != "manylinux": + continue + if jax_ref not in CI_JAX_REFS: + continue + + matrix.append( + { + "python_version": str(cell["python_version"]), + "jax_ref": jax_ref, + "jax_repository": str(cell["jax_repository"]), + "build_mode": build_mode, + "gfx_arch": str(cell["gfx_arch"]), + } + ) + + if not matrix: + raise ValueError( + "No supported manylinux JAX CI matrix entries were generated. " + f"Allowed refs: {sorted(CI_JAX_REFS)}" + ) + + return matrix + + def main(argv: list[str] | None = None) -> int: parser = argparse.ArgumentParser(description="Generate JAX release build matrix") parser.add_argument( diff --git a/build_tools/github_actions/configure_multi_arch_ci.py b/build_tools/github_actions/configure_multi_arch_ci.py index de9596d9bf2..1bc84a83b0c 100755 --- a/build_tools/github_actions/configure_multi_arch_ci.py +++ b/build_tools/github_actions/configure_multi_arch_ci.py @@ -66,7 +66,10 @@ get_git_submodule_paths, is_ci_run_required, ) -from configure_jax_release_matrix import generate_jax_matrix +from configure_jax_release_matrix import ( + generate_jax_matrix, + generate_jax_matrix_for_ci, +) from configure_pytorch_release_matrix import generate_pytorch_matrix_for_release_type from configure_rocm_python_test_matrix import build_rocm_python_test_matrix from github_actions_api import ( @@ -1065,7 +1068,13 @@ def _expand_build_config_for_platform( jax_build_matrix: list[dict[str, str]] = [] build_jax = jobs.build_jax.action == JobAction.RUN and platform == "linux" if build_jax: - jax_build_matrix = generate_jax_matrix(ci_inputs.python_versions or None) + requested_python_versions = ci_inputs.python_versions or None + + if ci_inputs.release_type == "ci": + jax_build_matrix = generate_jax_matrix_for_ci(requested_python_versions) + else: + jax_build_matrix = generate_jax_matrix(requested_python_versions) + # Flip back to False if the generated matrix is empty. build_jax = bool(jax_build_matrix) From 031550f82c42f0d1ad3cac0825b277e11ee09702 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 07:25:01 +0000 Subject: [PATCH 02/14] Add multi_arch_build_linux_jax_wheels_ci file name into _GITHUB_WORKFLOWS_CI_FILENAMES --- build_tools/github_actions/configure_ci_path_filters.py | 1 + 1 file changed, 1 insertion(+) diff --git a/build_tools/github_actions/configure_ci_path_filters.py b/build_tools/github_actions/configure_ci_path_filters.py index 376d75fd4fe..4ff666188d9 100644 --- a/build_tools/github_actions/configure_ci_path_filters.py +++ b/build_tools/github_actions/configure_ci_path_filters.py @@ -246,6 +246,7 @@ def is_ci_run_required(paths: Optional[Iterable[str]]) -> bool: "multi_arch_build_native_linux_packages.yml", "multi_arch_build_portable_linux_artifacts.yml", "multi_arch_build_portable_linux_pytorch_wheels_ci.yml", + "multi_arch_build_linux_jax_wheels_ci.yml", "multi_arch_build_portable_linux.yml", "multi_arch_build_windows_artifacts.yml", "multi_arch_build_windows_pytorch_wheels_ci.yml", From 4751f9ee7ddcfa66272f4d176cf834eb27419239 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 07:44:03 +0000 Subject: [PATCH 03/14] Include JAX fields in Linux BuildConfig contract --- .../github_actions/tests/configure_multi_arch_ci_test.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/build_tools/github_actions/tests/configure_multi_arch_ci_test.py b/build_tools/github_actions/tests/configure_multi_arch_ci_test.py index 88b9b2bfb4c..f0d37042bbb 100644 --- a/build_tools/github_actions/tests/configure_multi_arch_ci_test.py +++ b/build_tools/github_actions/tests/configure_multi_arch_ci_test.py @@ -1455,15 +1455,12 @@ def test_linux_workflow_uses_all_ci_fields(self): workflow_path = WORKFLOWS_DIR / "multi_arch_ci_linux.yml" yaml_fields = self._extract_build_config_fields(workflow_path) python_fields = {f.name for f in fields(cm.BuildConfig)} - # JAX builds are release-only for now, so CI workflows do not consume - # the JAX matrix fields even though setup reports them in summaries. - release_only_fields = {"build_jax", "jax_build_matrix"} self.assertEqual( yaml_fields, - python_fields - release_only_fields, + python_fields, f"BuildConfig fields mismatch with {workflow_path.name}.\n" f" In YAML but not Python: {yaml_fields - python_fields}\n" - f" In Python but not YAML: {python_fields - yaml_fields - release_only_fields}", + f" In Python but not YAML: {python_fields - yaml_fields}", ) def test_windows_workflow_uses_all_ci_fields(self): From f69fc011b73b4d20086dfd9cffb8e0174dd688e6 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 11:17:15 +0000 Subject: [PATCH 04/14] Remove github.ref_name --- .github/workflows/multi_arch_build_linux_jax_wheels_ci.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml index 141fea57724..f533e303c83 100644 --- a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml +++ b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml @@ -119,7 +119,7 @@ permissions: jobs: build_jax_wheels: - name: Build | py ${{ inputs.python_version }} | jax ${{ inputs.jax_ref }} + name: Build JAX | py ${{ inputs.python_version }} | jax ${{ inputs.jax_ref }} runs-on: ${{ github.repository_owner == 'ROCm' && 'azure-linux-scale-rocm' || 'ubuntu-24.04' }} outputs: package_find_links_url: ${{ steps.upload_jax_wheels.outputs.package_find_links_url }} @@ -136,7 +136,7 @@ jobs: uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 with: repository: ${{ inputs.repository || github.repository }} - ref: ${{ inputs.ref || github.ref_name }} + ref: ${{ inputs.ref }} - name: Checkout rocm-jax uses: actions/checkout@9c091bb21b7c1c1d1991bb908d89e4e9dddfe3e0 # v7.0.0 From c36e23fd2381a0a3c6a052902dc78f42284fe17d Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 19:04:32 +0000 Subject: [PATCH 05/14] Limit Python matrix for only Python 3.12 --- .github/workflows/multi_arch_build_linux_jax_wheels_ci.yml | 2 +- build_tools/github_actions/configure_jax_release_matrix.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml index f533e303c83..92ceb000606 100644 --- a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml +++ b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml @@ -119,7 +119,7 @@ permissions: jobs: build_jax_wheels: - name: Build JAX | py ${{ inputs.python_version }} | jax ${{ inputs.jax_ref }} + name: Build JAX | multi-arch-release | jax ${{ inputs.jax_ref }} | py ${{ inputs.python_version }} runs-on: ${{ github.repository_owner == 'ROCm' && 'azure-linux-scale-rocm' || 'ubuntu-24.04' }} outputs: package_find_links_url: ${{ steps.upload_jax_wheels.outputs.package_find_links_url }} diff --git a/build_tools/github_actions/configure_jax_release_matrix.py b/build_tools/github_actions/configure_jax_release_matrix.py index 2633c61a83c..7e63dcccf82 100644 --- a/build_tools/github_actions/configure_jax_release_matrix.py +++ b/build_tools/github_actions/configure_jax_release_matrix.py @@ -37,7 +37,7 @@ ] CI_JAX_REFS = {"rocm-jaxlib-v0.10.0", "rocm-jaxlib-v0.10.2"} -CI_PYTHON_VERSIONS = PYTHON_VERSIONS +CI_PYTHON_VERSIONS = ["3.12"] def _split_values(raw: str) -> list[str]: From 4bfebedd9afd7a333cb57653212f3fa95ba2dd88 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 21:25:56 +0000 Subject: [PATCH 06/14] keep only rocm-jaxlib-v0.10.0 version --- build_tools/github_actions/configure_jax_release_matrix.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/build_tools/github_actions/configure_jax_release_matrix.py b/build_tools/github_actions/configure_jax_release_matrix.py index 7e63dcccf82..d8b9ce35a7f 100644 --- a/build_tools/github_actions/configure_jax_release_matrix.py +++ b/build_tools/github_actions/configure_jax_release_matrix.py @@ -36,7 +36,7 @@ }, ] -CI_JAX_REFS = {"rocm-jaxlib-v0.10.0", "rocm-jaxlib-v0.10.2"} +CI_JAX_REFS = {"rocm-jaxlib-v0.10.0"} CI_PYTHON_VERSIONS = ["3.12"] From 3c1e222187add164a2e422e63729557ff4b810f6 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Thu, 9 Jul 2026 21:41:00 +0000 Subject: [PATCH 07/14] Adjust min matrix size to CI --- .../github_actions/tests/configure_multi_arch_ci_test.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/build_tools/github_actions/tests/configure_multi_arch_ci_test.py b/build_tools/github_actions/tests/configure_multi_arch_ci_test.py index f0d37042bbb..5c812f72ef2 100644 --- a/build_tools/github_actions/tests/configure_multi_arch_ci_test.py +++ b/build_tools/github_actions/tests/configure_multi_arch_ci_test.py @@ -1128,7 +1128,7 @@ def test_build_config_includes_jax_build_matrix(self): ) self.assertTrue(result.linux.build_jax) - self.assertGreater(len(result.linux.jax_build_matrix), 1) + self.assertGreater(len(result.linux.jax_build_matrix), 0) self.assertEqual( {row["python_version"] for row in result.linux.jax_build_matrix}, {"3.12"}, From 282210d2361491683b1872cf6bd6ff5288c1f439 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 04:03:41 +0000 Subject: [PATCH 08/14] Match PyTorch release-type matrix pattern --- .../configure_jax_release_matrix.py | 187 ++++++++++++------ .../github_actions/configure_multi_arch_ci.py | 16 +- 2 files changed, 133 insertions(+), 70 deletions(-) diff --git a/build_tools/github_actions/configure_jax_release_matrix.py b/build_tools/github_actions/configure_jax_release_matrix.py index d8b9ce35a7f..cf54ec7806c 100644 --- a/build_tools/github_actions/configure_jax_release_matrix.py +++ b/build_tools/github_actions/configure_jax_release_matrix.py @@ -14,30 +14,56 @@ from github_actions.github_actions_api import gha_set_output -PYTHON_VERSIONS = ["3.11", "3.12", "3.13", "3.14"] -JAX_REFS = [ - { +RELEASE_TYPES = ["ci", "dev", "nightly", "prerelease"] + +# TODO: add opt-ins for CI runs to use python versions and JAX refs normally +# only included in release runs. +RELEASE_PYTHON_VERSIONS = ["3.11", "3.12", "3.13", "3.14"] +CI_PYTHON_VERSIONS = { + "linux": ["3.12"], +} + +JAX_REF_CONFIGS = { + "rocm-jaxlib-v0.9.1": { "jax_ref": "rocm-jaxlib-v0.9.1", "jax_repository": "ROCm/rocm-jax", "build_mode": "native", "gfx_arch": "", }, - { + "rocm-jaxlib-v0.10.0": { "jax_ref": "rocm-jaxlib-v0.10.0", "jax_repository": "ROCm/jax", "build_mode": "manylinux", "gfx_arch": "device-all", }, - { + "rocm-jaxlib-v0.10.2": { "jax_ref": "rocm-jaxlib-v0.10.2", "jax_repository": "ROCm/jax", "build_mode": "manylinux", "gfx_arch": "device-all", }, -] - -CI_JAX_REFS = {"rocm-jaxlib-v0.10.0"} -CI_PYTHON_VERSIONS = ["3.12"] +} + +# Keep release behavior equivalent to the old generate_jax_matrix(None): +# all release refs across all release Python versions. +# +# TODO: separate out nightly/dev/prerelease JAX refs if those release types +# should differ later. +RELEASE_JAX_REFS = { + "linux": [ + "rocm-jaxlib-v0.9.1", + "rocm-jaxlib-v0.10.0", + "rocm-jaxlib-v0.10.2", + ], +} + +# CI uses the manylinux package path only. Exclude rocm-jaxlib-v0.9.1 because +# it uses the legacy native/tarball flow. +CI_JAX_REFS = { + "linux": [ + "rocm-jaxlib-v0.10.0", + ], +} def _split_values(raw: str) -> list[str]: @@ -49,78 +75,119 @@ def _split_values(raw: str) -> list[str]: ] -def generate_jax_matrix( - python_versions: list[str] | None, -) -> list[dict[str, object]]: - versions = python_versions if python_versions else PYTHON_VERSIONS - matrix: list[dict[str, object]] = [] - for py in versions: - for ref_cfg in JAX_REFS: - matrix.append( - { - "python_version": py, - "jax_ref": ref_cfg["jax_ref"], - "jax_repository": ref_cfg["jax_repository"], - "build_mode": ref_cfg["build_mode"], - "gfx_arch": ref_cfg["gfx_arch"], - } - ) - return matrix +def _default_python_versions(*, release_type: str, platform: str) -> list[str]: + if release_type == "ci": + return list(CI_PYTHON_VERSIONS[platform]) + return list(RELEASE_PYTHON_VERSIONS) -def generate_jax_matrix_for_ci( - python_versions: list[str] | None, +def _default_jax_refs(*, release_type: str, platform: str) -> list[str]: + if release_type == "ci": + return list(CI_JAX_REFS[platform]) + return list(RELEASE_JAX_REFS[platform]) + + +def generate_jax_matrix_for_release_type( + *, + release_type: str, + platform: str, + python_versions: list[str] | None = None, + jax_refs: list[str] | None = None, ) -> list[dict[str, str]]: - """Generate the JAX matrix used by Multi-Arch CI. + if release_type not in RELEASE_TYPES: + raise ValueError(f"Unknown release_type: {release_type!r}") + if platform not in ["linux"]: + raise ValueError(f"Unknown platform: {platform!r}") + + versions = python_versions or _default_python_versions( + release_type=release_type, + platform=platform, + ) + refs = jax_refs or _default_jax_refs( + release_type=release_type, + platform=platform, + ) - The release matrix includes legacy/native JAX entries such as - rocm-jaxlib-v0.9.1. The CI workflow is manylinux-only, so filter the - release matrix to the supported manylinux entries. - """ - versions = python_versions if python_versions else CI_PYTHON_VERSIONS + unknown_refs = sorted(set(refs) - set(JAX_REF_CONFIGS)) + if unknown_refs: + raise ValueError(f"Unknown JAX refs: {unknown_refs!r}") matrix: list[dict[str, str]] = [] - for cell in generate_jax_matrix(versions): - build_mode = str(cell["build_mode"]) - jax_ref = str(cell["jax_ref"]) - - if build_mode != "manylinux": - continue - if jax_ref not in CI_JAX_REFS: - continue - - matrix.append( - { - "python_version": str(cell["python_version"]), - "jax_ref": jax_ref, - "jax_repository": str(cell["jax_repository"]), - "build_mode": build_mode, - "gfx_arch": str(cell["gfx_arch"]), + for py in versions: + for ref in refs: + ref_cfg = JAX_REF_CONFIGS[ref] + row: dict[str, str] = { + "python_version": py, + "jax_ref": ref_cfg["jax_ref"], + "jax_repository": ref_cfg["jax_repository"], + "build_mode": ref_cfg["build_mode"], + "gfx_arch": ref_cfg["gfx_arch"], } - ) - - if not matrix: - raise ValueError( - "No supported manylinux JAX CI matrix entries were generated. " - f"Allowed refs: {sorted(CI_JAX_REFS)}" - ) + matrix.append(row) return matrix +def generate_jax_matrix( + python_versions: list[str] | None, +) -> list[dict[str, str]]: + """Generate the full JAX release matrix. + + Kept as a compatibility wrapper for callers that still expect the old + release-only helper. + """ + return generate_jax_matrix_for_release_type( + release_type="dev", + platform="linux", + python_versions=python_versions, + ) + + def main(argv: list[str] | None = None) -> int: - parser = argparse.ArgumentParser(description="Generate JAX release build matrix") + parser = argparse.ArgumentParser(description="Generate JAX build matrix") parser.add_argument( "--python-versions", type=str, default="", - help="Comma, semicolon, or whitespace separated list of Python versions (default: all)", + help=( + "Comma, semicolon, or whitespace separated list of Python versions " + "(default depends on --release-type)" + ), + ) + parser.add_argument( + "--jax-refs", + type=str, + default="", + help=( + "Comma, semicolon, or whitespace separated list of JAX refs " + "(default depends on --release-type and --platform)" + ), + ) + parser.add_argument( + "--platform", + type=str, + default="linux", + choices=["linux"], + help="Platform to generate matrix for (default: linux)", + ) + parser.add_argument( + "--release-type", + type=str, + default="dev", + choices=RELEASE_TYPES, + help="Release type selecting default JAX/Python matrix (default: dev)", ) args = parser.parse_args(argv) python_versions = _split_values(args.python_versions) or None + jax_refs = _split_values(args.jax_refs) or None - matrix = generate_jax_matrix(python_versions) + matrix = generate_jax_matrix_for_release_type( + release_type=args.release_type, + platform=args.platform, + python_versions=python_versions, + jax_refs=jax_refs, + ) gha_set_output({"jax_matrix": json.dumps(matrix)}) return 0 diff --git a/build_tools/github_actions/configure_multi_arch_ci.py b/build_tools/github_actions/configure_multi_arch_ci.py index 1bc84a83b0c..faa4c0ccfce 100755 --- a/build_tools/github_actions/configure_multi_arch_ci.py +++ b/build_tools/github_actions/configure_multi_arch_ci.py @@ -66,10 +66,7 @@ get_git_submodule_paths, is_ci_run_required, ) -from configure_jax_release_matrix import ( - generate_jax_matrix, - generate_jax_matrix_for_ci, -) +from configure_jax_release_matrix import generate_jax_matrix_for_release_type from configure_pytorch_release_matrix import generate_pytorch_matrix_for_release_type from configure_rocm_python_test_matrix import build_rocm_python_test_matrix from github_actions_api import ( @@ -1068,12 +1065,11 @@ def _expand_build_config_for_platform( jax_build_matrix: list[dict[str, str]] = [] build_jax = jobs.build_jax.action == JobAction.RUN and platform == "linux" if build_jax: - requested_python_versions = ci_inputs.python_versions or None - - if ci_inputs.release_type == "ci": - jax_build_matrix = generate_jax_matrix_for_ci(requested_python_versions) - else: - jax_build_matrix = generate_jax_matrix(requested_python_versions) + jax_build_matrix = generate_jax_matrix_for_release_type( + release_type=ci_inputs.release_type, + platform=platform, + python_versions=ci_inputs.python_versions or None, + ) # Flip back to False if the generated matrix is empty. build_jax = bool(jax_build_matrix) From f492fc2b929d9b6754ba422101b76672585841ea Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 04:39:44 +0000 Subject: [PATCH 09/14] Refactor configure_jax_release_matrix.py --- .../configure_jax_release_matrix.py | 52 +++++++++---------- 1 file changed, 24 insertions(+), 28 deletions(-) diff --git a/build_tools/github_actions/configure_jax_release_matrix.py b/build_tools/github_actions/configure_jax_release_matrix.py index cf54ec7806c..e201cc5d249 100644 --- a/build_tools/github_actions/configure_jax_release_matrix.py +++ b/build_tools/github_actions/configure_jax_release_matrix.py @@ -87,6 +87,27 @@ def _default_jax_refs(*, release_type: str, platform: str) -> list[str]: return list(RELEASE_JAX_REFS[platform]) +def generate_jax_matrix( + *, + jax_refs: list[str], + python_versions: list[str], +) -> list[dict[str, str]]: + matrix: list[dict[str, str]] = [] + for py in python_versions: + for ref in jax_refs: + ref_cfg = JAX_REF_CONFIGS[ref] + row: dict[str, str] = { + "python_version": py, + "jax_ref": ref_cfg["jax_ref"], + "jax_repository": ref_cfg["jax_repository"], + "build_mode": ref_cfg["build_mode"], + "gfx_arch": ref_cfg["gfx_arch"], + } + matrix.append(row) + + return matrix + + def generate_jax_matrix_for_release_type( *, release_type: str, @@ -112,34 +133,9 @@ def generate_jax_matrix_for_release_type( if unknown_refs: raise ValueError(f"Unknown JAX refs: {unknown_refs!r}") - matrix: list[dict[str, str]] = [] - for py in versions: - for ref in refs: - ref_cfg = JAX_REF_CONFIGS[ref] - row: dict[str, str] = { - "python_version": py, - "jax_ref": ref_cfg["jax_ref"], - "jax_repository": ref_cfg["jax_repository"], - "build_mode": ref_cfg["build_mode"], - "gfx_arch": ref_cfg["gfx_arch"], - } - matrix.append(row) - - return matrix - - -def generate_jax_matrix( - python_versions: list[str] | None, -) -> list[dict[str, str]]: - """Generate the full JAX release matrix. - - Kept as a compatibility wrapper for callers that still expect the old - release-only helper. - """ - return generate_jax_matrix_for_release_type( - release_type="dev", - platform="linux", - python_versions=python_versions, + return generate_jax_matrix( + jax_refs=refs, + python_versions=versions, ) From a36e6fe34b91515be06237d97e6087fabf5ad57f Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 04:53:12 +0000 Subject: [PATCH 10/14] Add unit tests --- .../configure_jax_release_matrix_test.py | 57 +++++++++++++++++-- 1 file changed, 53 insertions(+), 4 deletions(-) diff --git a/build_tools/github_actions/tests/configure_jax_release_matrix_test.py b/build_tools/github_actions/tests/configure_jax_release_matrix_test.py index ccabb23c403..d9d628e6448 100644 --- a/build_tools/github_actions/tests/configure_jax_release_matrix_test.py +++ b/build_tools/github_actions/tests/configure_jax_release_matrix_test.py @@ -13,8 +13,10 @@ class ConfigureJaxReleaseMatrixTest(unittest.TestCase): def test_default_matrix_includes_multiple_python_versions_and_refs(self): - matrix = m.generate_jax_matrix(None) - + matrix = m.generate_jax_matrix_for_release_type( + release_type="dev", + platform="linux", + ) python_versions = {row["python_version"] for row in matrix} jax_refs = {row["jax_ref"] for row in matrix} @@ -27,14 +29,61 @@ def test_default_matrix_includes_multiple_python_versions_and_refs(self): ) def test_explicit_python_version_narrows_matrix(self): - matrix = m.generate_jax_matrix(["3.12"]) - + matrix = m.generate_jax_matrix_for_release_type( + release_type="dev", + platform="linux", + python_versions=["3.12"], + ) self.assertGreater(len(matrix), 1) self.assertEqual( {row["python_version"] for row in matrix}, {"3.12"}, ) + def test_generate_jax_matrix_uses_requested_refs_only(self): + matrix = m.generate_jax_matrix( + jax_refs=["rocm-jaxlib-v0.10.0"], + python_versions=["3.12"], + ) + + self.assertEqual(len(matrix), 1) + self.assertEqual(matrix[0]["python_version"], "3.12") + self.assertEqual(matrix[0]["jax_ref"], "rocm-jaxlib-v0.10.0") + self.assertEqual(matrix[0]["jax_repository"], "ROCm/jax") + self.assertEqual(matrix[0]["build_mode"], "manylinux") + self.assertEqual(matrix[0]["gfx_arch"], "device-all") + + def test_ci_jax_matrix_excludes_unsupported_build_modes(self): + matrix = m.generate_jax_matrix_for_release_type( + release_type="ci", + platform="linux", + ) + + self.assertGreater(len(matrix), 0) + self.assertEqual( + {row["build_mode"] for row in matrix}, + {"manylinux"}, + ) + self.assertNotIn( + "rocm-jaxlib-v0.9.1", + {row["jax_ref"] for row in matrix}, + ) + + def test_unknown_release_type_raises(self): + with self.assertRaises(ValueError): + m.generate_jax_matrix_for_release_type( + release_type="unknown", + platform="linux", + ) + + def test_unknown_jax_ref_raises(self): + with self.assertRaises(ValueError): + m.generate_jax_matrix_for_release_type( + release_type="ci", + platform="linux", + jax_refs=["unknown-jax-ref"], + ) + if __name__ == "__main__": unittest.main() From 69d11b0bd320a8abeeb728b372e1af0a04df46b2 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 05:20:36 +0000 Subject: [PATCH 11/14] Set package dist dir in job env --- .github/workflows/multi_arch_build_linux_jax_wheels_ci.yml | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml index 92ceb000606..c1b0b62fff4 100644 --- a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml +++ b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml @@ -119,7 +119,7 @@ permissions: jobs: build_jax_wheels: - name: Build JAX | multi-arch-release | jax ${{ inputs.jax_ref }} | py ${{ inputs.python_version }} + name: Build JAX | ${{ inputs.artifact_group }} | jax ${{ inputs.jax_ref }} | py ${{ inputs.python_version }} runs-on: ${{ github.repository_owner == 'ROCm' && 'azure-linux-scale-rocm' || 'ubuntu-24.04' }} outputs: package_find_links_url: ${{ steps.upload_jax_wheels.outputs.package_find_links_url }} @@ -130,6 +130,7 @@ jobs: env: MANYLINUX_IMAGE_TAG: therock-jax-manylinux:${{ inputs.jax_ref }}-${{ inputs.python_version }} RELEASE_TYPE: ci + PACKAGE_DIST_DIR: ${{ github.workspace }}/jax-source/dist steps: - name: Checkout TheRock @@ -174,10 +175,6 @@ jobs: exit 1 fi - - name: Set package dist dir - run: | - echo "PACKAGE_DIST_DIR=${GITHUB_WORKSPACE}/jax-source/dist" >> "$GITHUB_ENV" - - name: Build manylinux image working-directory: jax-source env: From 33c156a42bf4133987176f9c59d7e03a42a88f14 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 06:01:05 +0000 Subject: [PATCH 12/14] Solve conflict --- .../configure_jax_release_matrix.py | 15 +++--- .../configure_jax_release_matrix_test.py | 53 ++++++++++++++++++- 2 files changed, 60 insertions(+), 8 deletions(-) diff --git a/build_tools/github_actions/configure_jax_release_matrix.py b/build_tools/github_actions/configure_jax_release_matrix.py index f89cccb423d..0f9ceb679f8 100644 --- a/build_tools/github_actions/configure_jax_release_matrix.py +++ b/build_tools/github_actions/configure_jax_release_matrix.py @@ -88,12 +88,14 @@ def _default_jax_refs(*, release_type: str, platform: str) -> list[str]: def generate_jax_matrix( - python_versions: list[str] | None, -) -> list[dict[str, object]]: - versions = python_versions if python_versions else PYTHON_VERSIONS - matrix: list[dict[str, object]] = [] - for py in versions: - for ref_cfg in JAX_REFS: + *, + jax_refs: list[str], + python_versions: list[str], +) -> list[dict[str, str]]: + matrix: list[dict[str, str]] = [] + for py in python_versions: + for ref in jax_refs: + ref_cfg = JAX_REF_CONFIGS[ref] # These row keys are the contract with workflow files which use them # via matrix. expressions. Empty values are allowed when the # workflow handles them explicitly, but undefined keys are not. @@ -109,6 +111,7 @@ def generate_jax_matrix( "gfx_arch": ref_cfg["gfx_arch"], } ) + return matrix diff --git a/build_tools/github_actions/tests/configure_jax_release_matrix_test.py b/build_tools/github_actions/tests/configure_jax_release_matrix_test.py index 8704328ffb1..cfc3f6374e3 100644 --- a/build_tools/github_actions/tests/configure_jax_release_matrix_test.py +++ b/build_tools/github_actions/tests/configure_jax_release_matrix_test.py @@ -23,6 +23,7 @@ def test_default_matrix_includes_multiple_python_versions_and_refs(self): release_type="dev", platform="linux", ) + python_versions = {row["python_version"] for row in matrix} jax_refs = {row["jax_ref"] for row in matrix} @@ -40,12 +41,42 @@ def test_explicit_python_version_narrows_matrix(self): platform="linux", python_versions=["3.12"], ) + self.assertGreater(len(matrix), 1) self.assertEqual( {row["python_version"] for row in matrix}, {"3.12"}, ) + def test_generate_jax_matrix_uses_requested_refs_only(self): + matrix = m.generate_jax_matrix( + jax_refs=["rocm-jaxlib-v0.10.0"], + python_versions=["3.12"], + ) + + self.assertEqual(len(matrix), 1) + self.assertEqual(matrix[0]["python_version"], "3.12") + self.assertEqual(matrix[0]["jax_ref"], "rocm-jaxlib-v0.10.0") + self.assertEqual(matrix[0]["jax_repository"], "ROCm/jax") + self.assertEqual(matrix[0]["build_mode"], "manylinux") + self.assertEqual(matrix[0]["gfx_arch"], "device-all") + + def test_ci_jax_matrix_excludes_unsupported_build_modes(self): + matrix = m.generate_jax_matrix_for_release_type( + release_type="ci", + platform="linux", + ) + + self.assertGreater(len(matrix), 0) + self.assertEqual( + {row["build_mode"] for row in matrix}, + {"manylinux"}, + ) + self.assertNotIn( + "rocm-jaxlib-v0.9.1", + {row["jax_ref"] for row in matrix}, + ) + def test_generated_rows_cover_workflow_matrix_inputs(self): # workflow file like: # @@ -62,14 +93,17 @@ def test_generated_rows_cover_workflow_matrix_inputs(self): # every row in the generated matrix. It intentionally does not check # that every generated key is consumed by each workflow; if we want to # enforce exact schemas, do that with generator-local tests. - workflow = load_workflow( WORKFLOWS_DIR / "multi_arch_release_linux_jax_wheels.yml" ) job = get_workflow_job(workflow, "build_jax_wheels") matrix_references = get_matrix_references(job["with"]) - matrix = m.generate_jax_matrix(["3.12"]) + matrix = m.generate_jax_matrix_for_release_type( + release_type="dev", + platform="linux", + python_versions=["3.12"], + ) self.assertGreater(len(matrix), 0) for row in matrix: @@ -79,6 +113,21 @@ def test_generated_rows_cover_workflow_matrix_inputs(self): # this test fails until the generator emits that key for every row. self.assertEqual(matrix_references - set(row), set()) + def test_unknown_release_type_raises(self): + with self.assertRaises(ValueError): + m.generate_jax_matrix_for_release_type( + release_type="unknown", + platform="linux", + ) + + def test_unknown_jax_ref_raises(self): + with self.assertRaises(ValueError): + m.generate_jax_matrix_for_release_type( + release_type="ci", + platform="linux", + jax_refs=["unknown-jax-ref"], + ) + if __name__ == "__main__": unittest.main() From 24b0eeb4a281d2c0143faa760ed9330c11c7317a Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 06:23:11 +0000 Subject: [PATCH 13/14] Add TODO 6388 --- .github/workflows/multi_arch_build_linux_jax_wheels_ci.yml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml index c1b0b62fff4..5ec8cb9be59 100644 --- a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml +++ b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml @@ -185,6 +185,10 @@ jobs: cp ../rocm-jax/docker/manylinux/Dockerfile.jax-manylinux_2_28-therock \ Dockerfile.jax-manylinux_2_28-therock + # TODO(#6388): Replace this per-run rocm-jax manylinux Dockerfile + # build with a frozen/prebuilt TheRock build image. Building this image + # in every CI run is expensive and currently rebuilds dependencies such + # as LLVM. docker build \ -t "${MANYLINUX_IMAGE_TAG}" \ --file=Dockerfile.jax-manylinux_2_28-therock \ From c4c61cb78916d13bbf96d2abca1568053698b7a7 Mon Sep 17 00:00:00 2001 From: erman-gurses Date: Wed, 15 Jul 2026 07:02:09 +0000 Subject: [PATCH 14/14] Update TODO note --- .github/workflows/multi_arch_build_linux_jax_wheels_ci.yml | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml index 5ec8cb9be59..5ef27e51927 100644 --- a/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml +++ b/.github/workflows/multi_arch_build_linux_jax_wheels_ci.yml @@ -185,10 +185,8 @@ jobs: cp ../rocm-jax/docker/manylinux/Dockerfile.jax-manylinux_2_28-therock \ Dockerfile.jax-manylinux_2_28-therock - # TODO(#6388): Replace this per-run rocm-jax manylinux Dockerfile - # build with a frozen/prebuilt TheRock build image. Building this image - # in every CI run is expensive and currently rebuilds dependencies such - # as LLVM. + # TODO(#6388): Replace this per-run rocm-jax Dockerfile build with a + # frozen/prebuilt CI image and refine when JAX CI should run. docker build \ -t "${MANYLINUX_IMAGE_TAG}" \ --file=Dockerfile.jax-manylinux_2_28-therock \