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
4 changes: 4 additions & 0 deletions docs/source/builder/build.md
Original file line number Diff line number Diff line change
Expand Up @@ -274,3 +274,7 @@ this check enabled, as it is one of the checks that validates that a kernel
is compliant. This option is primarily intended for kernels with
`triton.autotune` decorators, which can fail because there is no GPU available
in the build sandbox.

Generating a kernel's public API symbols (`symbols.json`) also requires
importing the kernel. So, when the `get_kernel` check is disabled, the build
variants will not contain a `symbols.json` file.
10 changes: 10 additions & 0 deletions docs/source/kernel-requirements.md
Original file line number Diff line number Diff line change
Expand Up @@ -486,6 +486,16 @@ __all__ = [
> [versioning guarantees](#versioning) apply to, so be sure to export
> every function, class, and `layers` module you want to expose.

The same applies to the `layers` module itself: only the layers listed in
the `__all__` of `layers` are part of the public API. For example:

```python
class SiluAndMul(nn.Module):
# ...

__all__ = ["SiluAndMul"]
```

## Python requirements

- Python code must be compatible with Python 3.9 and later.
Expand Down
77 changes: 72 additions & 5 deletions examples/kernels/flake.nix
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,12 @@
torchVersion = "213";
tvmFfiVersion = "01";

# Expected public symbols of the relu example kernels.
reluSymbols = {
functions = [ "relu" ];
layers = [ "ReLU" ];
};

# All example kernels to build in CI.
#
# - name: name in the output path
Expand All @@ -34,6 +40,9 @@
# - checkCudaCapabilities: optional list of CUDA capabilities (e.g. "9.0").
# When set, the kernel dylib must contain exactly this set of
# capabilities.
# - checkSymbols: optional attrset `{ functions = [ ... ]; layers = [ ... ]; }`.
# When set, every build variant must contain a symbols.json with
# exactly these public function and layer names.
ciKernels = [
{
name = "cpp20-symbols-kernel";
Expand All @@ -45,6 +54,7 @@
name = "relu-kernel";
path = ./relu;
drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${cudaVersion}-${sys}"};
checkSymbols = reluSymbols;
checkCudaCapabilities = [
"7.0"
"7.2"
Expand Down Expand Up @@ -97,6 +107,13 @@
path = ./relu-tvm-ffi;
drv =
sys: out: out.packages.${sys}.redistributable.${"tvm-ffi${tvmFfiVersion}-${cudaVersion}-${sys}"};
checkSymbols = {
functions = [
"relu"
"relu_jax"
];
layers = [ ];
};
}
{
name = "relu-tvm-ffi-compiler-flags-kernel";
Expand Down Expand Up @@ -293,6 +310,7 @@
name = "relu-triton-kernel";
path = ./relu-triton;
drv = sys: out: out.packages.${sys}.redistributable.torch-xpu;
checkSymbols = reluSymbols;
}
{
name = "gemm-triton-autotune-kernel";
Expand All @@ -303,12 +321,20 @@
name = "relu-kernel";
path = ./relu;
drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-${xpuVersion}-${sys}"};
checkSymbols = reluSymbols;
}
{
name = "relu-tvm-ffi-kernel";
path = ./relu-tvm-ffi;
drv =
sys: out: out.packages.${sys}.redistributable.${"tvm-ffi${tvmFfiVersion}-${xpuVersion}-${sys}"};
checkSymbols = {
functions = [
"relu"
"relu_jax"
];
layers = [ ];
};
}
{
name = "relu-tvm-ffi-compiler-flags-kernel";
Expand Down Expand Up @@ -345,6 +371,7 @@
name = "relu-kernel";
path = ./relu;
drv = sys: out: out.packages.${sys}.redistributable.${"torch${torchVersion}-metal-${sys}"};
checkSymbols = reluSymbols;
}
{
name = "relu-metal-cpp-kernel";
Expand Down Expand Up @@ -456,6 +483,41 @@
ln -s ${drv} $out
'';

# Check that every build variant in the output of `drv` contains a
# symbols.json with exactly the expected public function and layer
# names. On success, the output is a symlink to `drv`.
checkSymbols =
drv: expected:
pkgs.runCommand "${drv.name}-check-symbols"
{
nativeBuildInputs = [ pkgs.jq ];
expected = builtins.toJSON expected;
passAsFile = [ "expected" ];
}
''
variants=$(find ${drv}/ -mindepth 2 -maxdepth 2 -name metadata.json)
if [ -z "$variants" ]; then
echo "no build variants found in ${drv}" >&2
exit 1
fi

jq -S . "$expectedPath" > expected
for metadata in $variants; do
symbols="$(dirname "$metadata")/symbols.json"
if [ ! -f "$symbols" ]; then
echo "missing $symbols" >&2
exit 1
fi
jq -S '{functions: [.functions[].name], layers: [.layers[].name]}' "$symbols" > actual
if ! diff -u expected actual; then
echo "unexpected symbols in $symbols" >&2
exit 1
fi
done

ln -s ${drv} $out
'';

checkRocmArchs =
drv: expectedArchs:
checkKernelArchs drv expectedArchs
Expand Down Expand Up @@ -484,18 +546,23 @@
drv =
let
baseDrv = kernel.drv system kernel.outputs;
archsDrv =
if kernel ? checkRocmArchs then
checkRocmArchs baseDrv kernel.checkRocmArchs
else if kernel ? checkCudaCapabilities then
checkCudaCapabilities baseDrv kernel.checkCudaCapabilities
else
baseDrv;
in
if kernel.assertFail or false then
pkgs.testers.testBuildFailure' {
drv = baseDrv;
expectedBuilderLogEntries = kernel.assertFailLogs or [ ];
}
else if kernel ? checkRocmArchs then
checkRocmArchs baseDrv kernel.checkRocmArchs
else if kernel ? checkCudaCapabilities then
checkCudaCapabilities baseDrv kernel.checkCudaCapabilities
else if kernel ? checkSymbols then
checkSymbols archsDrv kernel.checkSymbols
else
baseDrv;
archsDrv;
}) kernelOutputsList;

mkCiBuild =
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,3 +7,6 @@
class ReLU(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return relu(x)


__all__ = ["ReLU"]
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
ops.relu(out, x)
return out


__all__ = ["ReLU"]
Original file line number Diff line number Diff line change
Expand Up @@ -7,3 +7,6 @@
class ReLU(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
return relu(x)


__all__ = ["ReLU"]
3 changes: 3 additions & 0 deletions examples/kernels/relu-triton/torch-ext/relu_triton/layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
ops.relu(out, x)
return out


__all__ = ["ReLU"]
Original file line number Diff line number Diff line change
Expand Up @@ -10,5 +10,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
relu(x, out)
return out

__all__ = ["ReLU"]
except ImportError:
pass
__all__ = []
Original file line number Diff line number Diff line change
Expand Up @@ -10,5 +10,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
relu(x, out)
return out

__all__ = ["ReLU"]
except ImportError:
pass
__all__ = []
3 changes: 3 additions & 0 deletions examples/kernels/relu/torch-ext/relu/layers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,3 +9,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
out = torch.empty_like(x)
ops.relu(out, x)
return out


__all__ = ["ReLU"]
5 changes: 5 additions & 0 deletions nix-builder/lib/extension/torch/arch.nix
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
cmake,
cmakeNvccThreadsHook,
cuda_nvcc,
generate-symbols-hook,
get-kernel-check,
hash-kernel-hook,
kernel-layout-check,
Expand Down Expand Up @@ -188,6 +189,10 @@ stdenv.mkDerivation (prevAttrs: {
python3 = python3.withPackages (ps: dependencies);
kernels = overrideTorch python3.pkgs.kernels;
})
(generate-symbols-hook.override {
python3 = python3.withPackages (ps: dependencies);
kernels = overrideTorch python3.pkgs.kernels;
})
]
++ lib.optionals cudaSupport [
cmakeNvccThreadsHook
Expand Down
5 changes: 5 additions & 0 deletions nix-builder/lib/extension/torch/no-arch.nix
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
stdenv,

kernel-builder,
generate-symbols-hook,
get-kernel-check,
hash-kernel-hook,
kernel-layout-check,
Expand Down Expand Up @@ -116,6 +117,10 @@ stdenv.mkDerivation (prevAttrs: {
python3 = python3.withPackages (_: dependencies);
kernels = overrideTorch python3.pkgs.kernels;
})
(generate-symbols-hook.override {
python3 = python3.withPackages (_: dependencies);
kernels = overrideTorch python3.pkgs.kernels;
})
];

buildPhase = ''
Expand Down
54 changes: 31 additions & 23 deletions nix-builder/lib/extension/tvm-ffi/arch.nix
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
cmake,
cmakeNvccThreadsHook,
cuda_nvcc,
generate-symbols-hook,
get-kernel-check,
hash-kernel-hook,
kernel-layout-check,
Expand Down Expand Up @@ -136,6 +137,30 @@ let

rustSupport = cargoLock != null;

# rpaths are stripped from kernels to make them portable, but that
# also means that in a Nix environment the CUDA/oneAPI dependencies
# cannot be located anymore, so pass them to hooks that load the kernel.
libraryPath = lib.makeLibraryPath (
map lib.getLib (
lib.optionals cudaSupport (
with cudaPackages;
[
cuda_cudart
libcublas
libcusolver
libcusparse
]
)
++ lib.optionals xpuSupport (
with xpuPackages;
[
intel-oneapi-compiler-dpcpp-cpp-runtime
intel-oneapi-compiler-shared-runtime
]
)
)
);

provenanceFlags = import ../provenance-flags.nix { inherit lib kernelProvenance; };

in
Expand Down Expand Up @@ -203,29 +228,12 @@ stdenv.mkDerivation (
(get-kernel-check.override {
python3 = python3.withPackages (ps: dependencies);
kernels = python3.pkgs.kernels.override { withTorch = false; };
# rpaths are stripped from kernels to make them portable, but that
# also means that in a Nix environment the CUDA/oneAPI dependencies
# cannot be located anymore, so pass them to get-kernel-check.
libraryPath = lib.makeLibraryPath (
map lib.getLib (
lib.optionals cudaSupport (
with cudaPackages;
[
cuda_cudart
libcublas
libcusolver
libcusparse
]
)
++ lib.optionals xpuSupport (
with xpuPackages;
[
intel-oneapi-compiler-dpcpp-cpp-runtime
intel-oneapi-compiler-shared-runtime
]
)
)
);
inherit libraryPath;
})
(generate-symbols-hook.override {
python3 = python3.withPackages (ps: dependencies);
kernels = python3.pkgs.kernels.override { withTorch = false; };
inherit libraryPath;
})
]
++ lib.optionals cudaSupport [
Expand Down
4 changes: 4 additions & 0 deletions nix-builder/overlay.nix
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,10 @@ final: prev:

fetchKernelDeps = final.callPackage ./pkgs/fetch-kernel-deps { };

generate-symbols = final.python3.pkgs.callPackage ./pkgs/generate-symbols { };

generate-symbols-hook = final.callPackage ./pkgs/generate-symbols/hook.nix { };

get-kernel-check = final.callPackage ./pkgs/get-kernel-check { };

hash-kernel-hook = final.callPackage ./pkgs/hash-kernel-hook { };
Expand Down
40 changes: 40 additions & 0 deletions nix-builder/pkgs/generate-symbols/default.nix
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
{
lib,
buildPythonPackage,
pytestCheckHook,
setuptools,

kernels,
}:

let
version = (builtins.fromTOML (builtins.readFile ./pyproject.toml)).project.version;
in
buildPythonPackage {
pname = "generate-symbols";
inherit version;
pyproject = true;

src = lib.fileset.toSource {
root = ./.;
fileset = lib.fileset.unions [
./pyproject.toml
./src
./tests
];
};

build-system = [ setuptools ];

dependencies = [ kernels ];

nativeCheckInputs = [ pytestCheckHook ];

pythonImportsCheck = [ "generate_symbols" ];

meta = {
description = "Generate public API symbols for built kernels";
license = lib.licenses.asl20;
mainProgram = "generate-symbols";
};
}
Loading
Loading