Skip to content

Commit 8465cee

Browse files
authored
feat: prefer aotrion from release (#541)
* feat: prefer aotrion from release * fix: resolve libamdhip * fix: remove installCheckPhase
1 parent 4a2d92f commit 8465cee

2 files changed

Lines changed: 36 additions & 166 deletions

File tree

Lines changed: 8 additions & 60 deletions
Original file line numberDiff line numberDiff line change
@@ -1,21 +1,11 @@
11
{
22
callPackage,
3-
fetchFromGitHub,
4-
fetchpatch,
53
fetchurl,
64
stdenvNoCC,
75
}:
86

97
let
108
generic = callPackage ./generic.nix { };
11-
postFetch = ''
12-
cd $out
13-
git reset --hard HEAD
14-
for submodule in $(git config --file .gitmodules --get-regexp path | awk '{print $2}' | grep '^third_party/' | grep -v '^third_party/triton$'); do
15-
git submodule update --init --recursive "$submodule"
16-
done
17-
find "$out" -name .git -print0 | xargs -0 rm -rf
18-
'';
199
mkImages =
2010
version: srcs:
2111
stdenvNoCC.mkDerivation {
@@ -35,31 +25,12 @@ in
3525
aotriton_0_11_1 = generic rec {
3626
version = "0.11.1b";
3727

38-
src = fetchFromGitHub {
39-
owner = "ROCm";
40-
repo = "aotriton";
41-
tag = version;
42-
hash = "sha256-F7JjyS+6gMdCpOFLldTsNJdVzzVwd6lwW7+V8ZOZfig=";
43-
leaveDotGit = true;
44-
inherit postFetch;
28+
hashes = {
29+
"7.0" = "sha256-3rgEbp75dsJzn9BWO1AjnhLcAC19T5fBxKGHSstlq8Q=";
30+
"7.1" = "sha256-wWE+2enuzHNZ8EoWJLtSjlT15jaeaC3URuqpNtlFI1g=";
31+
"7.2" = "sha256-VsoxJUwWVfpNUWji2zFZeBwkQtX2sBiCFV6ThZuFzxY=";
4532
};
4633

47-
patches = [
48-
# Fails with: ld.lld: error: unable to insert .comment after .comment
49-
./v0.11.1b-no-ld-script.diff
50-
];
51-
52-
gpuTargets = [
53-
# aotriton GPU support list:
54-
# https://github.com/ROCm/aotriton/blob/main/v2python/gpu_targets.py
55-
"gfx90a"
56-
"gfx942"
57-
"gfx950"
58-
"gfx1100"
59-
"gfx1151"
60-
"gfx1201"
61-
];
62-
6334
images = mkImages version [
6435
(fetchurl {
6536
url = "https://github.com/ROCm/aotriton/releases/download/0.11.1b/aotriton-0.11.1b-images-amd-gfx90a.tar.gz";
@@ -82,38 +53,17 @@ in
8253
hash = "sha256-Ck/zJL/9rAwv3oeop/cFY9PISoCtTo8xNF8rQKE4TpU=";
8354
})
8455
];
85-
86-
extraPythonDepends = ps: [ ps.pandas ];
8756
};
8857

8958
aotriton_0_11_2 = generic rec {
9059
version = "0.11.2b";
9160

92-
src = fetchFromGitHub {
93-
owner = "ROCm";
94-
repo = "aotriton";
95-
tag = version;
96-
hash = "sha256-VIwwQR1fl40NLNOwO8KhQK/xOK6wb2l8qBugJ1cRjm4=";
97-
leaveDotGit = true;
98-
inherit postFetch;
61+
hashes = {
62+
"7.0" = "sha256-VQGgo7MAiQABtmJfKjU5p7rWDzhvCgYevn1O1coPr7k=";
63+
"7.1" = "sha256-/uNr6z6khM4YFVu6/gJsV3/WcF5EaeWUBbJgvXS4zBA=";
64+
"7.2" = "sha256-zYq/J7u2POxFyUE16bKHRZZgdCY6awVV5YeK4ctqI0k=";
9965
};
10066

101-
patches = [
102-
# Fails with: ld.lld: error: unable to insert .comment after .comment
103-
./v0.11.1b-no-ld-script.diff
104-
];
105-
106-
gpuTargets = [
107-
# aotriton GPU support list:
108-
# https://github.com/ROCm/aotriton/blob/main/v2python/gpu_targets.py
109-
"gfx90a"
110-
"gfx942"
111-
"gfx950"
112-
"gfx1100"
113-
"gfx1151"
114-
"gfx1201"
115-
];
116-
11767
images = mkImages version [
11868
(fetchurl {
11969
url = "https://github.com/ROCm/aotriton/releases/download/0.11.2b/aotriton-0.11.2b-images-amd-gfx90a.tar.gz";
@@ -136,8 +86,6 @@ in
13686
hash = "sha256-Ck/zJL/9rAwv3oeop/cFY9PISoCtTo8xNF8rQKE4TpU=";
13787
})
13888
];
139-
140-
extraPythonDepends = ps: [ ps.pandas ];
14189
};
14290

14391
}
Lines changed: 28 additions & 106 deletions
Original file line numberDiff line numberDiff line change
@@ -1,137 +1,59 @@
1-
# Vendored from nixpkgs
21
{
2+
autoPatchelfHook,
3+
clr,
4+
fetchurl,
35
lib,
6+
rocm-core,
47
stdenv,
5-
cmake,
6-
jq,
7-
python3,
8-
ninja,
9-
pkg-config,
10-
rocmPackages,
11-
writableTmpDirAsHomeHook,
12-
writeShellScriptBin,
138
xz,
149
}:
1510

1611
{
1712
version,
18-
gpuTargets,
19-
patches ? [ ],
20-
src,
2113
images,
22-
extraPythonDepends ? ps: [ ],
14+
hashes,
2315
}:
2416

2517
let
26-
gpuTargets' = lib.concatStringsSep ";" gpuTargets;
27-
compiler = "amdclang++";
18+
rocmVersion = lib.versions.majorMinor rocm-core.version;
19+
hash =
20+
hashes.${rocmVersion}
21+
or (throw "aotriton ${version} binary package is not specified for ROCm ${rocmVersion}");
2822
in
29-
stdenv.mkDerivation (finalAttrs: {
23+
stdenv.mkDerivation {
3024
pname = "aotriton";
25+
inherit version;
3126

32-
inherit version src patches;
33-
34-
env = {
35-
#CXX = compiler;
36-
ROCM_PATH = "${rocmPackages.clr}";
37-
CFLAGS = "-w -g1 -gz -Wno-c++11-narrowing";
38-
CXXFLAGS = finalAttrs.env.CFLAGS;
39-
40-
# aotriton passes a lot of files to the linker.
41-
NIX_LD_USE_RESPONSE_FILE = 1;
27+
src = fetchurl {
28+
url = "https://github.com/ROCm/aotriton/releases/download/${version}/aotriton-${version}-manylinux_2_28_x86_64-rocm${rocmVersion}-shared.tar.gz";
29+
inherit hash;
4230
};
4331

44-
requiredSystemFeatures = [ "big-parallel" ];
45-
46-
nativeBuildInputs = [
47-
cmake
48-
jq
49-
rocmPackages.rocm-cmake
50-
pkg-config
51-
python3
52-
ninja
53-
rocmPackages.clr
54-
writableTmpDirAsHomeHook # venv wants to cache in ~
55-
(writeShellScriptBin "amdclang++" ''
56-
exec ${rocmPackages.llvm.clang}/bin/clang++ "$@"
57-
'')
58-
];
59-
32+
nativeBuildInputs = [ autoPatchelfHook ];
6033
buildInputs = [
61-
rocmPackages.clr
34+
clr
35+
stdenv.cc.cc.lib
6236
xz
63-
]
64-
++ (with python3.pkgs; [
65-
wheel
66-
packaging
67-
pyyaml
68-
numpy
69-
filelock
70-
iniconfig
71-
pluggy
72-
pybind11
73-
pandas
74-
triton
75-
]);
76-
77-
preConfigure = lib.optionalString (lib.versionAtLeast version "0.11.1") ''
78-
# Since we use pre-built images, we can grab the image SHA from there.
79-
# As of 0.11.1b this doesn't seem to be used for image loading yet, but
80-
# just in case this happens in the future, we set this to the actual
81-
# value and not a stub.
82-
export AOTRITON_CI_SUPPLIED_SHA1=$(jq -r '.["AOTRITON_GIT_SHA1"]' ${images}/lib/aotriton.images/amd-gfx90a/__signature__)
83-
84-
# Need to set absolute paths to VENV and its PYTHON or
85-
# build fails with "AOTRITON_INHERIT_SYSTEM_SITE_TRITON is enabled
86-
# but triton is not available … no such file or directory"
87-
# Set via a preConfigure hook so a valid absolute path can be
88-
# picked if nix-shell is used against this package
89-
cmakeFlagsArray+=(
90-
"-DVENV_DIR=$(pwd)/aotriton-venv/"
91-
"-DVENV_BIN_PYTHON=$(pwd)/aotriton-venv/bin/python"
92-
)
93-
'';
37+
];
9438

95-
# From README:
96-
# Note: do not run ninja separately, due to the limit of the current build system,
97-
# ninja install will run the whole build process unconditionally.
39+
dontConfigure = true;
9840
dontBuild = true;
99-
41+
dontStrip = true;
10042
installPhase = ''
10143
runHook preInstall
102-
ninja -v install
103-
ln -sf ${images}/lib/aotriton.images $out/lib/aotriton.images
104-
runHook postInstall
105-
'';
10644
107-
doCheck = false;
108-
doInstallCheck = false;
45+
mkdir -p "$out"
46+
cp -r include lib "$out/"
47+
ln -s ${images}/lib/aotriton.images "$out/lib/aotriton.images"
10948
110-
cmakeFlags = [
111-
# Disable building kernels if no supported targets are enabled
112-
(lib.cmakeBool "AOTRITON_NOIMAGE_MODE" true)
113-
# Use preinstalled triton from our python's site-packages
114-
(lib.cmakeBool "AOTRITON_INHERIT_SYSTEM_SITE_TRITON" true)
115-
# Avoid kernels being skipped if build host is overloaded
116-
(lib.cmakeFeature "AOTRITON_GPU_BUILD_TIMEOUT" "0")
117-
(lib.cmakeFeature "CMAKE_CXX_COMPILER" compiler)
118-
# Manually define CMAKE_INSTALL_<DIR>
119-
# See: https://github.com/NixOS/nixpkgs/pull/197838
120-
(lib.cmakeFeature "CMAKE_INSTALL_BINDIR" "bin")
121-
(lib.cmakeFeature "CMAKE_INSTALL_LIBDIR" "lib")
122-
(lib.cmakeFeature "CMAKE_INSTALL_INCLUDEDIR" "include")
123-
(lib.cmakeFeature "AOTRITON_TARGET_ARCH" gpuTargets')
124-
(lib.cmakeBool "AOTRITON_USE_TORCH" false)
125-
];
49+
runHook postInstall
50+
'';
12651

12752
meta = with lib; {
12853
description = "Ahead of Time (AOT) Triton Math Library";
12954
homepage = "https://github.com/ROCm/aotriton";
13055
license = with licenses; [ mit ];
131-
platforms = platforms.linux;
132-
sourceProvenance = with sourceTypes; [
133-
fromSource
134-
binaryNativeCode # aotriton.images
135-
];
56+
platforms = [ "x86_64-linux" ];
57+
sourceProvenance = with sourceTypes; [ binaryNativeCode ];
13658
};
137-
})
59+
}

0 commit comments

Comments
 (0)