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
3 changes: 3 additions & 0 deletions docs/source/builder/build-variants.md
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,9 @@ available. This list will be updated as new PyTorch versions are released.

## XPU x86_64-linux

- `torch210-xpu20253-x86_64-linux`
- `torch211-xpu20253-x86_64-linux`
- `torch212-xpu20253-x86_64-linux`
- `torch213-xpu20260-x86_64-linux`
- `torch214-xpu20261-x86_64-linux`

Expand Down
11 changes: 11 additions & 0 deletions examples/kernels/flake.nix
Original file line number Diff line number Diff line change
Expand Up @@ -326,6 +326,17 @@
path = ./cutlass-gemm;
drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${xpuVersion}-${sys}"};
}
{
name = "relu-torch-stable-abi-kernel";
path = ./relu-torch-stable-abi;
drv = sys: out: out.packages.${sys}.redistributable.${"torch-stable-abi211-${xpuVersion}-${sys}"};
}
# Check that we can also build the non-stable ABI version preceding the stable ABI version.
{
name = "relu-torch-stable-abi-kernel";
path = ./relu-torch-stable-abi;
drv = sys: out: out.packages.${sys}.redistributable.${"torch210-xpu20253-${sys}"};
}
];

# Metal kernels to build in CI.
Expand Down
12 changes: 5 additions & 7 deletions kernel-builder/src/pyproject/templates/torch/preamble.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -78,19 +78,17 @@ set(_STABLE_ABI_VERSION_{{ entry.backend }} "{{ entry.version }}")
{% endfor %}
set(_STABLE_ABI_VERSION "${_STABLE_ABI_VERSION_${BACKEND}}")

if(_STABLE_ABI_VERSION)
if (TORCH_VERSION VERSION_LESS ${_STABLE_ABI_VERSION})
message(FATAL_ERROR "Torch version ${TORCH_VERSION} is less than the stable ABI "
"version ${_STABLE_ABI_VERSION}. Cannot build with stable ABI targeting a newer version of Torch.")
endif()

if(_STABLE_ABI_VERSION AND TORCH_VERSION VERSION_GREATER_EQUAL ${_STABLE_ABI_VERSION})
message("Building for the Torch stable ABI. ABI version: ${_STABLE_ABI_VERSION}, Torch version: ${TORCH_VERSION}")
# From the Torch docs: TORCH_TARGET_VERSION (((0ULL + major) << 56) | ((0ULL + minor) << 48))
string(REPLACE "." ";" _STABLE_ABI_VERSION_LIST "${_STABLE_ABI_VERSION}")
list(GET _STABLE_ABI_VERSION_LIST 0 _STABLE_ABI_MAJOR)
list(GET _STABLE_ABI_VERSION_LIST 1 _STABLE_ABI_MINOR)
math(EXPR _STABLE_ABI_HEX "(${_STABLE_ABI_MAJOR} << 56) | (${_STABLE_ABI_MINOR} << 48)" OUTPUT_FORMAT HEXADECIMAL)

add_compile_definitions(-DTORCH_TARGET_VERSION=${_STABLE_ABI_HEX})
else()
message("Not building for the Torch stable ABI. ABI version: ${_STABLE_ABI_VERSION}, Torch version: ${TORCH_VERSION}")
unset(_STABLE_ABI_VERSION)
endif()
{% endif %}

Expand Down
3 changes: 3 additions & 0 deletions nix-builder/build-variants.json
Original file line number Diff line number Diff line change
Expand Up @@ -46,6 +46,9 @@
"torch214-rocm72-x86_64-linux"
],
"xpu": [
"torch210-xpu20253-x86_64-linux",
"torch211-xpu20253-x86_64-linux",
"torch212-xpu20253-x86_64-linux",
"torch213-xpu20260-x86_64-linux",
"torch214-xpu20261-x86_64-linux"
]
Expand Down
21 changes: 19 additions & 2 deletions nix-builder/lib/build.nix
Original file line number Diff line number Diff line change
Expand Up @@ -88,8 +88,18 @@ rec {
in
builtins.attrValues newestPerGroup;

# Split up build sets into two:
#
# - Right: build sets with Torch versions that can build for the
# kernels requested stable ABI version.
# - Wrong: all other build sets.
Comment thread
danieldk marked this conversation as resolved.
#
# For instance, if stable-abi = 2.13 for XPU, then all Torch >= 2.13
# build sets will be in `right` and all Torch < 2.13 build sets will
# be in `wrong`.
byStableAbi = lib.partition (
buildSet: kernelConfig.isTorchStableAbiForBackend buildSet.buildConfig.backend
buildSet:
kernelConfig.torchCoversStableAbi buildSet.buildConfig.backend buildSet.buildConfig.torchVersion
) (buildSetsWithinBounds buildSets);
in
deduplicateForStableAbi byStableAbi.right ++ byStableAbi.wrong;
Expand Down Expand Up @@ -180,6 +190,13 @@ rec {
variant = variants.kernelVariant kernelConfig;
}
else
let
torchStableAbiVersion =
if kernelConfig.torchCoversStableAbi buildConfig.backend buildConfig.torchVersion then
kernelConfig.torchStableAbiVersionForBackend buildConfig.backend
else
null;
in
extension.mkTorchExtension {
inherit
buildConfig
Expand All @@ -195,7 +212,7 @@ rec {
kernelProvenance
;

torchStableAbiVersion = kernelConfig.torchStableAbiVersionForBackend buildConfig.backend;
inherit torchStableAbiVersion;

kernelName = kernelConfig.name;
doAbiCheck = true;
Expand Down
8 changes: 8 additions & 0 deletions nix-builder/lib/kernel-config.nix
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,14 @@ in
# Does the given backend use the torch stable ABI.
isTorchStableAbiForBackend = backend: torchStableAbiVersionForBackend backend != null;

# The given Torch version can build for this kernel's ABI version.
torchCoversStableAbi =
backend: torchVersion:
let
stableAbiVersion = torchStableAbiVersionForBackend backend;
in
stableAbiVersion != null && lib.versionAtLeast torchVersion stableAbiVersion;

# Kernel backends.
backends =
let
Expand Down
4 changes: 3 additions & 1 deletion nix-builder/lib/variants/torch.nix
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,9 @@ in
archVariant = kernelConfig.kernelBackends.${buildConfig.backend};
stableAbiVersion = kernelConfig.torchStableAbiVersionForBackend buildConfig.backend;
in
if archVariant && stableAbiVersion != null then
if
archVariant && kernelConfig.torchCoversStableAbi buildConfig.backend buildConfig.torchVersion

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🧠 syntax. Not gonna lie.

then
"torch-stable-abi${flattenVersion (lib.versions.majorMinor stableAbiVersion)}-${computeString}-${buildConfig.system}"
else if archVariant then
archString
Expand Down
10 changes: 6 additions & 4 deletions nix-builder/overlay.nix
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,8 @@ final: prev:
;
inherit (triton-rocm) triton-rocm_3_7_0;
inherit (triton-xpu)
triton-xpu_3_6_0
triton-xpu_3_7_0
triton-xpu_3_7_1
triton-xpu_3_7_2
triton-xpu_3_8_0
Expand Down Expand Up @@ -238,17 +240,17 @@ final: prev:
version = "2.10";
triton-cuda = null;
triton-rocm = null;
triton-xpu = null;
xpuPackages = null;
triton-xpu = triton-xpu_3_6_0;
xpuPackages = final.xpuPackages_2025_3_1;
};

# Maintain a minimal version for TPU support.
torch-bin_2_11 = mkTorch {
version = "2.11";
triton-cuda = null;
triton-rocm = null;
triton-xpu = null;
xpuPackages = null;
triton-xpu = triton-xpu_3_7_0;
xpuPackages = final.xpuPackages_2025_3_2;
};

torch-bin_2_12 = mkTorch {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,13 @@
"hash": "sha256-yIsRKf1OFPD4gpY8ZygxXKrjXS9HN00X7e7R7cdpdJc=",
"version": "2.10.0"
}
},
"x86_64-linux": {
"xpu": {
"url": "https://download.pytorch.org/whl/xpu/torch-2.10.0%2Bxpu-cp314-cp314-linux_x86_64.whl",
"hash": "sha256-NJN2IhJKNBK+qpizLyYcHLVnZGEwcv/DQ/TaGzQg7X8=",
"version": "2.10.0"
}
}
},
"2.11": {
Expand All @@ -14,6 +21,11 @@
"url": "https://download.pytorch.org/whl/cpu/torch-2.11.0%2Bcpu-cp314-cp314-manylinux_2_28_x86_64.whl",
"hash": "sha256-+EgeqQiOTluBF4p1qr27ZYvehjm8GhX9XY+TCryWZzU=",
"version": "2.11.0"
},
"xpu": {
"url": "https://download.pytorch.org/whl/xpu/torch-2.11.0%2Bxpu-cp314-cp314-linux_x86_64.whl",
"hash": "sha256-9j+b6Gd+E4oquM6pYA4F1Db3AFvgEmr3Clj3s/cYnnM=",
"version": "2.11.0"
}
},
"aarch64-darwin": {
Expand All @@ -31,6 +43,13 @@
"hash": "sha256-99+uSlGRl9+gUOmNjjY3ig+1iZYlqHXCtURFAFouQE4=",
"version": "2.12.0"
}
},
"x86_64-linux": {
"xpu": {
"url": "https://download.pytorch.org/whl/xpu/torch-2.12.0%2Bxpu-cp314-cp314-linux_x86_64.whl",
"hash": "sha256-CAXv6u5uWyIA13Bx3WxS1lUg+VsD5H9S9oPmL0jmIuY=",
"version": "2.12.0"
}
}
},
"2.13": {
Expand Down
15 changes: 15 additions & 0 deletions nix-builder/pkgs/python-modules/torch/binary/torch-versions.json
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,11 @@
"metal": true,
"systems": ["aarch64-darwin"]
},
{
"torchVersion": "2.10.0",
"xpuVersion": "2025.3.1",
"systems": ["x86_64-linux"]
},

{
"torchVersion": "2.11.0",
Expand All @@ -16,12 +21,22 @@
"metal": true,
"systems": ["aarch64-darwin"]
},
{
"torchVersion": "2.11.0",
"xpuVersion": "2025.3.2",
"systems": ["x86_64-linux"]
},

{
"torchVersion": "2.12.0",
"metal": true,
"systems": ["aarch64-darwin"]
},
{
"torchVersion": "2.12.0",
"xpuVersion": "2025.3.2",
"systems": ["x86_64-linux"]
},

{
"torchVersion": "2.13.0",
Expand Down
12 changes: 12 additions & 0 deletions nix-builder/pkgs/python-modules/triton-xpu/default.nix
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,18 @@ let
generic = callPackage ./generic.nix { };
in
{
triton-xpu_3_6_0 = generic {
version = "3.6.0";
url = "https://download.pytorch.org/whl/triton_xpu-3.6.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl";
hash = "sha256-WM1zFm+4jIkNP35TyaslJUFhT6ygcJ6Ju15cYeNjwgw=";
};

triton-xpu_3_7_0 = generic {
version = "3.7.0";
url = "https://download.pytorch.org/whl/triton_xpu-3.7.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl";
hash = "sha256-HiBLsYRJjwIo3eC8dkeIrvSjE15SCejSWQJ1hmbDCqg=";
};

triton-xpu_3_7_1 = generic {
version = "3.7.1";
url = "https://download-r2.pytorch.org/whl/triton_xpu-3.7.1-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl";
Expand Down
Loading
Loading