Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

44 changes: 22 additions & 22 deletions docs/source/builder/build-variants.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
2 changes: 1 addition & 1 deletion docs/source/builder/build.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
8 changes: 4 additions & 4 deletions docs/source/builder/ide-setup.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
```

Expand Down Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion docs/source/builder/writing-kernels.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.<variant>`. 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
Expand Down
4 changes: 2 additions & 2 deletions docs/source/kernel-requirements.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
`<framework><version>-cxx<abiver>-<cu><cudaver>-<arch>-<os>`.
For example `build/torch26-cxx98-cu118-x86_64-linux`.
`<framework><version>-<cu><cudaver>-<arch>-<os>`.
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
Expand Down
48 changes: 17 additions & 31 deletions examples/kernels/flake.nix
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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"
Expand All @@ -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"
Expand All @@ -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";
Expand All @@ -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";
Expand All @@ -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" ];
}
Expand Down Expand Up @@ -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"
Expand All @@ -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"
Expand All @@ -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}"};
}
];

Expand Down Expand Up @@ -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";
Expand All @@ -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}"};
}
];

Expand All @@ -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};
Expand Down
19 changes: 7 additions & 12 deletions kernel-builder/src/pyproject/templates/torch/build-variants.cmake
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
# Generate a standardized build variant name following the pattern:
# torch<VERSION>-[cxx11-]<COMPUTE>-<ARCH>-<OS>
# torch<VERSION>-<COMPUTE>-<ARCH>-<OS>
# or, when compiled against the Torch stable ABI:
# torch-stable-abi<VERSION>-<COMPUTE>-<ARCH>-<OS>
#
Expand All @@ -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<VERSION> (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)
Expand Down Expand Up @@ -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}")
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading