diff --git a/Cargo.lock b/Cargo.lock index 7cc80f36e..dd9ffae11 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1871,7 +1871,7 @@ version = "7.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5f98efec8807c63c752b5bd61f862c165c115b0a35685bdcfd9238c7aeb592b7" dependencies = [ - "cfg-if 1.0.4", + "cfg-if", "unicode-width 0.1.14", ] diff --git a/docs/source/builder/build-variants.md b/docs/source/builder/build-variants.md index c41f038c4..5f3226c68 100644 --- a/docs/source/builder/build-variants.md +++ b/docs/source/builder/build-variants.md @@ -17,43 +17,43 @@ available. This list will be updated as new PyTorch versions are released. ## CPU aarch64-linux -- `torch213-cxx11-cpu-aarch64-linux` -- `torch214-cxx11-cpu-aarch64-linux` +- `torch213-cpu-aarch64-linux` +- `torch214-cpu-aarch64-linux` ## CUDA aarch64-linux -- `torch213-cxx11-cu126-aarch64-linux` -- `torch213-cxx11-cu130-aarch64-linux` -- `torch213-cxx11-cu132-aarch64-linux` -- `torch214-cxx11-cu126-aarch64-linux` -- `torch214-cxx11-cu130-aarch64-linux` -- `torch214-cxx11-cu132-aarch64-linux` +- `torch213-cu126-aarch64-linux` +- `torch213-cu130-aarch64-linux` +- `torch213-cu132-aarch64-linux` +- `torch214-cu126-aarch64-linux` +- `torch214-cu130-aarch64-linux` +- `torch214-cu132-aarch64-linux` ## CPU x86_64-linux -- `torch213-cxx11-cpu-x86_64-linux` -- `torch214-cxx11-cpu-x86_64-linux` +- `torch213-cpu-x86_64-linux` +- `torch214-cpu-x86_64-linux` ## CUDA x86_64-linux -- `torch213-cxx11-cu126-x86_64-linux` -- `torch213-cxx11-cu130-x86_64-linux` -- `torch213-cxx11-cu132-x86_64-linux` -- `torch214-cxx11-cu126-x86_64-linux` -- `torch214-cxx11-cu130-x86_64-linux` -- `torch214-cxx11-cu132-x86_64-linux` +- `torch213-cu126-x86_64-linux` +- `torch213-cu130-x86_64-linux` +- `torch213-cu132-x86_64-linux` +- `torch214-cu126-x86_64-linux` +- `torch214-cu130-x86_64-linux` +- `torch214-cu132-x86_64-linux` ## ROCm x86_64-linux -- `torch213-cxx11-rocm71-x86_64-linux` -- `torch213-cxx11-rocm72-x86_64-linux` -- `torch214-cxx11-rocm714-x86_64-linux` -- `torch214-cxx11-rocm72-x86_64-linux` +- `torch213-rocm71-x86_64-linux` +- `torch213-rocm72-x86_64-linux` +- `torch214-rocm714-x86_64-linux` +- `torch214-rocm72-x86_64-linux` ## XPU x86_64-linux -- `torch213-cxx11-xpu20260-x86_64-linux` -- `torch214-cxx11-xpu20261-x86_64-linux` +- `torch213-xpu20260-x86_64-linux` +- `torch214-xpu20261-x86_64-linux` ## Python-only kernels diff --git a/docs/source/builder/build.md b/docs/source/builder/build.md index 124d2870e..71beb224a 100644 --- a/docs/source/builder/build.md +++ b/docs/source/builder/build.md @@ -98,7 +98,7 @@ using: ```bash $ rm -rf .venv # Remove existing venv if any. -$ kernel-builder devshell --variant torch212-cxx11-rocm71-x86_64-linux +$ kernel-builder devshell --variant torch214-rocm72-x86_64-linux ``` For an editor-driven workflow with `direnv` activating the devshell on diff --git a/docs/source/builder/ide-setup.md b/docs/source/builder/ide-setup.md index a689b1844..f22f07d89 100644 --- a/docs/source/builder/ide-setup.md +++ b/docs/source/builder/ide-setup.md @@ -88,7 +88,7 @@ the `Creating new venv environment in path: './.venv'` line from the To pin a non-default build variant, name it explicitly: ```bash -$ echo 'use flake .#devShells.torch212-cxx11-rocm71-x86_64-linux' > .envrc +$ echo 'use flake .#devShells.torch214-rocm71-x86_64-linux' > .envrc $ direnv allow ``` @@ -182,13 +182,13 @@ variant. For example: ```bash # CUDA 13.0 -use flake .#devShells.torch212-cxx11-cu130-x86_64-linux +use flake .#devShells.torch214-cu130-x86_64-linux # ROCm 7.1 -use flake .#devShells.torch212-cxx11-rocm71-x86_64-linux +use flake .#devShells.torch214-rocm72-x86_64-linux # XPU -use flake .#devShells.torch212-cxx11-xpu20253-x86_64-linux +use flake .#devShells.torch212-xpu20261-x86_64-linux ``` Remove `.venv/` first if it was created against a different variant, diff --git a/docs/source/builder/writing-kernels.md b/docs/source/builder/writing-kernels.md index b381f842b..a3e25b3b7 100644 --- a/docs/source/builder/writing-kernels.md +++ b/docs/source/builder/writing-kernels.md @@ -539,7 +539,7 @@ ROCm, XPU, Metal, CPU. If you would like to the tests for a specific build variant, you can use `nix run .#ciTests.`. For instance: ```bash -$ nix run .#ciTests.torch210-cxx11-cpu-x86_64-linux +$ nix run .#ciTests.torch210-cpu-x86_64-linux ``` When running the tests on a non-NixOS systems, make sure that diff --git a/docs/source/kernel-requirements.md b/docs/source/kernel-requirements.md index 67fbd72ef..ee3074c9c 100644 --- a/docs/source/kernel-requirements.md +++ b/docs/source/kernel-requirements.md @@ -46,8 +46,8 @@ displays a corresponding badge in the UI. A kernel repository on the Hub must contain a `build` directory. This directory contains build variants of a kernel in the form of directories following the template -`-cxx---`. -For example `build/torch26-cxx98-cu118-x86_64-linux`. +`---`. +For example `build/torch214-cu130-x86_64-linux`. The kernel is in the build variant directory and must contain a `__init__.py` file. For compatibility with older versions of the diff --git a/examples/kernels/flake.nix b/examples/kernels/flake.nix index 553ad75c1..7acc3fffd 100644 --- a/examples/kernels/flake.nix +++ b/examples/kernels/flake.nix @@ -38,14 +38,13 @@ { name = "cpp20-symbols-kernel"; path = ./cpp20-symbols; - drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-cpu-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cpu-${sys}"}; } # This test should check the capabilities of the oldest supported CUDA. { name = "relu-kernel"; path = ./relu; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cudaVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"}; checkCudaCapabilities = [ "7.0" "7.2" @@ -62,8 +61,7 @@ name = "relu-kernel-cu13"; path = ./relu; drv = - sys: out: - out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cuda13Version}-${sys}"}; + sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cuda13Version}-${sys}"}; checkCudaCapabilities = [ "7.5" "8.0" @@ -82,8 +80,7 @@ # Check arch intersection. 5.0 is dropped because it is not supported. name = "relu-archs-subset"; path = ./relu-archs-subset; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cudaVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"}; checkCudaCapabilities = [ "7.5" "8.0" @@ -110,13 +107,12 @@ { name = "relu-kernel-cpu"; path = ./relu; - drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-cpu-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cpu-${sys}"}; } { name = "cutlass-gemm-kernel"; path = ./cutlass-gemm; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cudaVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"}; } { name = "cutlass-gemm-tvm-ffi-kernel"; @@ -127,8 +123,7 @@ { name = "relu-backprop-compile-kernel"; path = ./relu-backprop-compile; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cudaVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"}; } { name = "silu-and-mul-kernel"; @@ -155,14 +150,12 @@ { name = "relu-compiler-flags"; path = ./relu-compiler-flags; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cudaVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"}; } { name = "relu-invalid-capability"; path = ./relu-invalid-capability; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${cudaVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"}; assertFail = true; assertFailLogs = [ "empty set of capabilities" ]; } @@ -223,16 +216,14 @@ { name = "relu-invalid-capability"; path = ./relu-invalid-capability; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${rocmVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${rocmVersion}-${sys}"}; assertFail = true; assertFailLogs = [ "empty set of architectures" ]; } { name = "relu-kernel"; path = ./relu; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${rocmVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${rocmVersion}-${sys}"}; checkRocmArchs = [ "gfx906" "gfx908" @@ -254,8 +245,7 @@ # Check arch intersection. gfx940 is dropped because it is not supported. name = "relu-archs-subset"; path = ./relu-archs-subset; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${rocmVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${rocmVersion}-${sys}"}; checkRocmArchs = [ "gfx90a" "gfx942" @@ -265,8 +255,7 @@ { name = "relu-compiler-flags"; path = ./relu-compiler-flags; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${rocmVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${rocmVersion}-${sys}"}; } ]; @@ -313,8 +302,7 @@ { name = "relu-kernel"; path = ./relu; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${xpuVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${xpuVersion}-${sys}"}; } { name = "relu-tvm-ffi-kernel"; @@ -331,14 +319,12 @@ { name = "relu-compiler-flags"; path = ./relu-compiler-flags; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${xpuVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${xpuVersion}-${sys}"}; } { name = "cutlass-gemm-kernel"; path = ./cutlass-gemm; - drv = - sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-cxx11-${xpuVersion}-${sys}"}; + drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${xpuVersion}-${sys}"}; } ]; @@ -365,7 +351,7 @@ drv = sys: _out: let - variant = "torch${torchVersion}-cxx11-cpu-${sys}"; + variant = "torch${torchVersion}-cpu-${sys}"; conflictsFlake = mkKernelOutputs { path = ./symbol-conflicts; }; conflicts2Flake = mkKernelOutputs { path = ./symbol-conflicts2; }; conflicts = conflictsFlake.packages.${sys}.redistributable.${variant}; diff --git a/kernel-builder/src/pyproject/templates/torch/build-variants.cmake b/kernel-builder/src/pyproject/templates/torch/build-variants.cmake index 46526fa94..06f6bb7d7 100644 --- a/kernel-builder/src/pyproject/templates/torch/build-variants.cmake +++ b/kernel-builder/src/pyproject/templates/torch/build-variants.cmake @@ -1,5 +1,5 @@ # Generate a standardized build variant name following the pattern: -# torch-[cxx11-]-- +# torch--- # or, when compiled against the Torch stable ABI: # torch-stable-abi--- # @@ -13,9 +13,9 @@ # TORCH_STABLE_ABI - Stable ABI version the extension was compiled against (e.g., "2.11"); # when set, TORCH_VERSION is ignored and the prefix becomes # torch-stable-abi (e.g., "2.11" -> "torch-stable-abi211") -# Example output: torch27-cxx11-cu124-x86_64-linux (Linux) -# torch27-cu124-x86_64-windows (Windows) -# torch27-metal-aarch64-darwin (macOS) +# Example output: torch214-cu124-x86_64-linux (Linux) +# torch214-cu124-x86_64-windows (Windows) +# torch214-metal-aarch64-darwin (macOS) # torch-stable-abi211-cu124-x86_64-linux (Linux, stable ABI) # function(generate_build_name OUT_BUILD_NAME TORCH_VERSION COMPUTE_FRAMEWORK COMPUTE_VERSION) @@ -117,12 +117,7 @@ function(generate_build_name OUT_BUILD_NAME TORCH_VERSION COMPUTE_FRAMEWORK COMP set(ARCH_OS_STRING "${CPU_ARCH}-${OS_NAME}") # Assemble the final build name - # For non-stable-ABI Linux builds, include cxx11 ABI indicator for compatibility - if(NOT ARG_TORCH_STABLE_ABI AND ARCH_OS_STRING MATCHES "-linux$") - set(BUILD_NAME "${TORCH_PREFIX}-cxx11-${COMPUTE_STRING}-${ARCH_OS_STRING}") - else() - set(BUILD_NAME "${TORCH_PREFIX}-${COMPUTE_STRING}-${ARCH_OS_STRING}") - endif() + set(BUILD_NAME "${TORCH_PREFIX}-${COMPUTE_STRING}-${ARCH_OS_STRING}") set(${OUT_BUILD_NAME} "${BUILD_NAME}" PARENT_SCOPE) message(STATUS "Generated build name: ${BUILD_NAME}") @@ -136,7 +131,7 @@ endfunction() # Arguments: # TARGET_NAME - Name of the target to create the install rule for # PACKAGE_NAME - Python package name (e.g., "activation") -# BUILD_VARIANT_NAME - Build variant name (e.g., "torch271-cxx11-cu124-x86_64-linux") +# BUILD_VARIANT_NAME - Build variant name (e.g., "torch214-cu124-x86_64-linux") # INSTALL_PREFIX - Base installation directory (defaults to CMAKE_INSTALL_PREFIX) # GPU_ARCHS - List of GPU architectures that were compiled # (optional; when provided for CUDA/ROCm, metadata.json will include @@ -234,7 +229,7 @@ endfunction() # Arguments: # TARGET_NAME - Name of the target to create the install rule for # PACKAGE_NAME - Python package name (e.g., "activation") -# BUILD_VARIANT_NAME - Build variant name (e.g., "torch271-cxx11-cu124-x86_64-linux") +# BUILD_VARIANT_NAME - Build variant name (e.g., "torch214-cu124-x86_64-linux") # GPU_ARCHS - List of GPU architectures that were compiled # (optional; when provided for CUDA/ROCm, metadata.json will include # a "backend" key with the type and arch list) diff --git a/kernels/tests/test_variants.py b/kernels/tests/test_variants.py index 427735549..fe2c6da2f 100644 --- a/kernels/tests/test_variants.py +++ b/kernels/tests/test_variants.py @@ -4,6 +4,7 @@ from kernels.backends import CPU, CUDA, Metal, ROCm from kernels.variants import ( + Variant, VariantAccepted, VariantRejected, _resolve_variant_for_system, @@ -150,26 +151,36 @@ def test_get_variants(): assert variant_strs.issuperset(SUPERSET_VARIANT_STRINGS) -RESOLVE_VARIANTS = [ - parse_variant(s) - for s in [ - "torch210-cxx11-cu128-x86_64-linux", - "torch210-cxx11-cu126-x86_64-linux", - "torch210-cxx11-cu130-x86_64-linux", - "torch210-cxx11-rocm70-x86_64-linux", - "torch210-cxx11-cpu-x86_64-linux", - "torch210-cpu-aarch64-darwin", - "torch210-metal-aarch64-darwin", - "torch-cuda", - "torch-cpu", +@pytest.fixture(params=["", "-cxx11"], ids=["tagless", "cxx11"]) +def linux_abi(request) -> str: + # Test build variants with and without the C++ ABI tag. kernel-builder + # used to add an ABI tag to distingiush between the C++98 and C++11 + # ABIs. + return request.param + + +def _resolve_variants(linux_abi: str) -> list[Variant]: + return [ + parse_variant(s) + for s in [ + f"torch210{linux_abi}-cu128-x86_64-linux", + f"torch210{linux_abi}-cu126-x86_64-linux", + f"torch210{linux_abi}-cu130-x86_64-linux", + f"torch210{linux_abi}-rocm70-x86_64-linux", + f"torch210{linux_abi}-cpu-x86_64-linux", + "torch210-cpu-aarch64-darwin", + "torch210-metal-aarch64-darwin", + "torch-cuda", + "torch-cpu", + ] ] -] -def test_resolve_cuda_exact(): +def test_resolve_cuda_exact(linux_abi): # CUDA 12.8 should resolve to cu128. + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CUDA(Version("12.8")), cpu="x86_64", os="linux", @@ -178,15 +189,16 @@ def test_resolve_cuda_exact(): tvm_ffi_version=None, ) assert result != [] - assert result[0].variant_str == "torch210-cxx11-cu128-x86_64-linux" + assert result[0].variant_str == f"torch210{linux_abi}-cu128-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_cuda_best_older_minor(): +def test_resolve_cuda_best_older_minor(linux_abi): # CUDA 12.9 is not available, should fall back to cu128 (highest <= 12.9). + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CUDA(Version("12.9")), cpu="x86_64", os="linux", @@ -195,15 +207,16 @@ def test_resolve_cuda_best_older_minor(): tvm_ffi_version=None, ) assert result != [] - assert result[0].variant_str == "torch210-cxx11-cu128-x86_64-linux" + assert result[0].variant_str == f"torch210{linux_abi}-cu128-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_cuda_no_newer_minor(): +def test_resolve_cuda_no_newer_minor(linux_abi): # CUDA 12.5 is older than all the variants, fall back to noarch. + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CUDA(Version("12.5")), cpu="x86_64", os="linux", @@ -214,13 +227,14 @@ def test_resolve_cuda_no_newer_minor(): assert result != [] assert result[0].variant_str == "torch-cuda" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_cuda_no_different_major(): +def test_resolve_cuda_no_different_major(linux_abi): # Different major version must not match. + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CUDA(Version("11.8")), cpu="x86_64", os="linux", @@ -231,12 +245,13 @@ def test_resolve_cuda_no_different_major(): assert result != [] assert result[0].variant_str == "torch-cuda" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_rocm(): +def test_resolve_rocm(linux_abi): + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=ROCm(Version("7.0")), cpu="x86_64", os="linux", @@ -245,14 +260,15 @@ def test_resolve_rocm(): tvm_ffi_version=None, ) assert result != [] - assert result[0].variant_str == "torch210-cxx11-rocm70-x86_64-linux" + assert result[0].variant_str == f"torch210{linux_abi}-rocm70-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_cpu_linux(): +def test_resolve_cpu_linux(linux_abi): + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CPU(), cpu="x86_64", os="linux", @@ -261,14 +277,15 @@ def test_resolve_cpu_linux(): tvm_ffi_version=None, ) assert result != [] - assert result[0].variant_str == "torch210-cxx11-cpu-x86_64-linux" + assert result[0].variant_str == f"torch210{linux_abi}-cpu-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_cpu_darwin(): +def test_resolve_cpu_darwin(linux_abi): + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CPU(), cpu="aarch64", os="darwin", @@ -279,12 +296,13 @@ def test_resolve_cpu_darwin(): assert result != [] assert result[0].variant_str == "torch210-cpu-aarch64-darwin" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_metal_darwin(): +def test_resolve_metal_darwin(linux_abi): + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=Metal(), cpu="aarch64", os="darwin", @@ -296,7 +314,7 @@ def test_resolve_metal_darwin(): assert result != [] assert result[0].variant_str == "torch210-metal-aarch64-darwin" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) RESOLVE_VARIANTS_METAL = [ @@ -349,10 +367,11 @@ def test_resolve_metal_darwin_new_macos(): assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_METAL) -def test_resolve_noarch_fallback(): +def test_resolve_noarch_fallback(linux_abi): # With no matching arch variant, should fall back to torch noarch. + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=CUDA(Version("12.8")), cpu="aarch64", os="linux", @@ -363,12 +382,13 @@ def test_resolve_noarch_fallback(): assert result != [] assert result[0].variant_str == "torch-cuda" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_no_match(): +def test_resolve_no_match(linux_abi): + variants = _resolve_variants(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS, + variants=variants, selected_backend=ROCm(Version("7.0")), cpu="x86_64", os="linux", @@ -378,22 +398,24 @@ def test_resolve_no_match(): ) assert result == [] assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS) + assert {vs.variant for vs in trace} == set(variants) -RESOLVE_VARIANTS_UNIVERSAL = [ - parse_variant(s) - for s in [ - "torch210-cxx11-cu128-x86_64-linux", - "torch-universal", +def _resolve_variants_universal(linux_abi: str) -> list[Variant]: + return [ + parse_variant(s) + for s in [ + f"torch210{linux_abi}-cu128-x86_64-linux", + "torch-universal", + ] ] -] -def test_resolve_universal_matches_any_backend(): +def test_resolve_universal_matches_any_backend(linux_abi): # Universal works with every backend. + variants = _resolve_variants_universal(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_UNIVERSAL, + variants=variants, selected_backend=ROCm(Version("7.0")), cpu="x86_64", os="linux", @@ -404,13 +426,14 @@ def test_resolve_universal_matches_any_backend(): assert result != [] assert result[0].variant_str == "torch-universal" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_UNIVERSAL) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_universal_is_last_resort(): +def test_resolve_universal_is_last_resort(linux_abi): # Specific match is preferred over universal. + variants = _resolve_variants_universal(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_UNIVERSAL, + variants=variants, selected_backend=CUDA(Version("12.8")), cpu="x86_64", os="linux", @@ -419,9 +442,9 @@ def test_resolve_universal_is_last_resort(): tvm_ffi_version=None, ) assert result != [] - assert result[0].variant_str == "torch210-cxx11-cu128-x86_64-linux" + assert result[0].variant_str == f"torch210{linux_abi}-cu128-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_UNIVERSAL) + assert {vs.variant for vs in trace} == set(variants) def test_resolve_specific_noarch_preferred_over_universal(): @@ -442,20 +465,22 @@ def test_resolve_specific_noarch_preferred_over_universal(): assert {vs.variant for vs in trace} == set(variants) -RESOLVE_VARIANTS_NO_NOARCH = [ - parse_variant(s) - for s in [ - "torch210-cxx11-cu126-x86_64-linux", - "torch210-cxx11-cu128-x86_64-linux", - "torch210-cxx11-cu130-x86_64-linux", +def _resolve_variants_no_noarch(linux_abi: str) -> list[Variant]: + return [ + parse_variant(s) + for s in [ + f"torch210{linux_abi}-cu126-x86_64-linux", + f"torch210{linux_abi}-cu128-x86_64-linux", + f"torch210{linux_abi}-cu130-x86_64-linux", + ] ] -] -def test_resolve_cuda_no_newer_minor_no_noarch(): +def test_resolve_cuda_no_newer_minor_no_noarch(linux_abi): # No compatible variant for 12.5. + variants = _resolve_variants_no_noarch(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_NO_NOARCH, + variants=variants, selected_backend=CUDA(Version("12.5")), cpu="x86_64", os="linux", @@ -465,13 +490,14 @@ def test_resolve_cuda_no_newer_minor_no_noarch(): ) assert result == [] assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_NO_NOARCH) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_cuda_no_different_major_no_noarch(): +def test_resolve_cuda_no_different_major_no_noarch(linux_abi): # 11.8 has a different major, so there is no compatible fallback. + variants = _resolve_variants_no_noarch(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_NO_NOARCH, + variants=variants, selected_backend=CUDA(Version("11.8")), cpu="x86_64", os="linux", @@ -481,23 +507,25 @@ def test_resolve_cuda_no_different_major_no_noarch(): ) assert result == [] assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_NO_NOARCH) + assert {vs.variant for vs in trace} == set(variants) -RESOLVE_VARIANTS_STABLE_ABI = [ - parse_variant(s) - for s in [ - "torch-stable-abi211-cu128-x86_64-linux", - "torch210-cxx11-cu128-x86_64-linux", - "torch-cuda", +def _resolve_variants_stable_abi(linux_abi: str) -> list[Variant]: + return [ + parse_variant(s) + for s in [ + "torch-stable-abi211-cu128-x86_64-linux", + f"torch210{linux_abi}-cu128-x86_64-linux", + "torch-cuda", + ] ] -] -def test_resolve_stable_abi_accepted(): +def test_resolve_stable_abi_accepted(linux_abi): # Stable ABI 2.11 is accepted when torch_version == stable ABI version. + variants = _resolve_variants_stable_abi(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_STABLE_ABI, + variants=variants, selected_backend=CUDA(Version("12.8")), cpu="x86_64", os="linux", @@ -508,13 +536,14 @@ def test_resolve_stable_abi_accepted(): assert result != [] assert result[0].variant_str == "torch-stable-abi211-cu128-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_STABLE_ABI) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_stable_abi_accepted_newer_torch(): +def test_resolve_stable_abi_accepted_newer_torch(linux_abi): # Stable ABI 2.11 is also accepted when torch_version > stable ABI version. + variants = _resolve_variants_stable_abi(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_STABLE_ABI, + variants=variants, selected_backend=CUDA(Version("12.8")), cpu="x86_64", os="linux", @@ -525,13 +554,14 @@ def test_resolve_stable_abi_accepted_newer_torch(): assert result != [] assert result[0].variant_str == "torch-stable-abi211-cu128-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_STABLE_ABI) + assert {vs.variant for vs in trace} == set(variants) -def test_resolve_stable_abi_rejected_newer_abi(): +def test_resolve_stable_abi_rejected_newer_abi(linux_abi): # Stable ABI 2.11 is rejected when torch_version < stable ABI version. + variants = _resolve_variants_stable_abi(linux_abi) result, trace = _resolve_variant_for_system( - variants=RESOLVE_VARIANTS_STABLE_ABI, + variants=variants, selected_backend=CUDA(Version("12.8")), cpu="x86_64", os="linux", @@ -540,9 +570,9 @@ def test_resolve_stable_abi_rejected_newer_abi(): tvm_ffi_version=None, ) assert result != [] - assert result[0].variant_str == "torch210-cxx11-cu128-x86_64-linux" + assert result[0].variant_str == f"torch210{linux_abi}-cu128-x86_64-linux" assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] - assert {vs.variant for vs in trace} == set(RESOLVE_VARIANTS_STABLE_ABI) + assert {vs.variant for vs in trace} == set(variants) def test_resolve_stable_abi_newest_version_preferred(): @@ -594,13 +624,43 @@ def test_resolve_tagless_preferred_over_abi_tagged(): assert {vs.variant for vs in trace} == set(variants) -def test_resolve_stable_abi_preferred_over_torch(): +def test_resolve_mixed_abi_tags(): + # The legacy cxx98 tag is rejected on a cxx11 Torch, while both the + # tagless and cxx11-tagged variants are accepted (tagless first). + variants = [ + parse_variant(s) + for s in [ + "torch210-cxx98-cu128-x86_64-linux", + "torch210-cxx11-cu128-x86_64-linux", + "torch210-cu128-x86_64-linux", + ] + ] + result, trace = _resolve_variant_for_system( + variants=variants, + selected_backend=CUDA(Version("12.8")), + cpu="x86_64", + os="linux", + torch_version=Version("2.10"), + torch_cxx11_abi=True, + tvm_ffi_version=None, + ) + assert [v.variant_str for v in result] == [ + "torch210-cu128-x86_64-linux", + "torch210-cxx11-cu128-x86_64-linux", + ] + rejected = {vs.variant.variant_str: vs.reason for vs in trace if isinstance(vs, VariantRejected)} + assert "CXX11 ABI" in rejected["torch210-cxx98-cu128-x86_64-linux"] + assert result == [vs.variant for vs in trace if isinstance(vs, VariantAccepted)] + assert {vs.variant for vs in trace} == set(variants) + + +def test_resolve_stable_abi_preferred_over_torch(linux_abi): # TorchStableAbi variant is preferred over a regular Torch variant of the same version. variants = [ parse_variant(s) for s in [ "torch-stable-abi211-cu128-x86_64-linux", - "torch211-cxx11-cu128-x86_64-linux", + f"torch211{linux_abi}-cu128-x86_64-linux", ] ] result, trace = _resolve_variant_for_system( diff --git a/nix-builder/build-variants.json b/nix-builder/build-variants.json index 9e07e6a36..bbef6fa9f 100644 --- a/nix-builder/build-variants.json +++ b/nix-builder/build-variants.json @@ -11,40 +11,40 @@ }, "aarch64-linux": { "cpu": [ - "torch213-cxx11-cpu-aarch64-linux", - "torch214-cxx11-cpu-aarch64-linux" + "torch213-cpu-aarch64-linux", + "torch214-cpu-aarch64-linux" ], "cuda": [ - "torch213-cxx11-cu126-aarch64-linux", - "torch213-cxx11-cu130-aarch64-linux", - "torch213-cxx11-cu132-aarch64-linux", - "torch214-cxx11-cu126-aarch64-linux", - "torch214-cxx11-cu130-aarch64-linux", - "torch214-cxx11-cu132-aarch64-linux" + "torch213-cu126-aarch64-linux", + "torch213-cu130-aarch64-linux", + "torch213-cu132-aarch64-linux", + "torch214-cu126-aarch64-linux", + "torch214-cu130-aarch64-linux", + "torch214-cu132-aarch64-linux" ] }, "x86_64-linux": { "cpu": [ - "torch213-cxx11-cpu-x86_64-linux", - "torch214-cxx11-cpu-x86_64-linux" + "torch213-cpu-x86_64-linux", + "torch214-cpu-x86_64-linux" ], "cuda": [ - "torch213-cxx11-cu126-x86_64-linux", - "torch213-cxx11-cu130-x86_64-linux", - "torch213-cxx11-cu132-x86_64-linux", - "torch214-cxx11-cu126-x86_64-linux", - "torch214-cxx11-cu130-x86_64-linux", - "torch214-cxx11-cu132-x86_64-linux" + "torch213-cu126-x86_64-linux", + "torch213-cu130-x86_64-linux", + "torch213-cu132-x86_64-linux", + "torch214-cu126-x86_64-linux", + "torch214-cu130-x86_64-linux", + "torch214-cu132-x86_64-linux" ], "rocm": [ - "torch213-cxx11-rocm71-x86_64-linux", - "torch213-cxx11-rocm72-x86_64-linux", - "torch214-cxx11-rocm714-x86_64-linux", - "torch214-cxx11-rocm72-x86_64-linux" + "torch213-rocm71-x86_64-linux", + "torch213-rocm72-x86_64-linux", + "torch214-rocm714-x86_64-linux", + "torch214-rocm72-x86_64-linux" ], "xpu": [ - "torch213-cxx11-xpu20260-x86_64-linux", - "torch214-cxx11-xpu20261-x86_64-linux" + "torch213-xpu20260-x86_64-linux", + "torch214-xpu20261-x86_64-linux" ] } } diff --git a/nix-builder/lib/build-variants.nix b/nix-builder/lib/build-variants.nix index 8948911f1..a18445d2d 100644 --- a/nix-builder/lib/build-variants.nix +++ b/nix-builder/lib/build-variants.nix @@ -32,7 +32,7 @@ rec { else if version.system == "aarch64-darwin" then "torch${flattenVersion version.torchVersion}-${computeString version}-${version.system}" else - "torch${flattenVersion version.torchVersion}-cxx11-${computeString version}-${version.system}"; + "torch${flattenVersion version.torchVersion}-${computeString version}-${version.system}"; # Build variants included in bundle builds. buildVariants = diff --git a/nix-builder/lib/variants/torch.nix b/nix-builder/lib/variants/torch.nix index 7b9469277..6ef189d10 100644 --- a/nix-builder/lib/variants/torch.nix +++ b/nix-builder/lib/variants/torch.nix @@ -24,7 +24,7 @@ let if buildConfig.system == "aarch64-darwin" then "${torchString}-${computeString}-${buildConfig.system}" else - "${torchString}-cxx11-${computeString}-${buildConfig.system}"; + "${torchString}-${computeString}-${buildConfig.system}"; in { arch = archString;