mirror of
https://github.com/NixOS/nixpkgs.git
synced 2026-09-28 10:50:15 +00:00
python3Packages.jax-cuda13-pjrt: init at 0.11.1
This commit is contained in:
132
pkgs/development/python-modules/jax-cuda13-pjrt/default.nix
Normal file
132
pkgs/development/python-modules/jax-cuda13-pjrt/default.nix
Normal file
@@ -0,0 +1,132 @@
|
||||
{
|
||||
lib,
|
||||
stdenv,
|
||||
buildPythonPackage,
|
||||
fetchPypi,
|
||||
addDriverRunpath,
|
||||
autoPatchelfHook,
|
||||
pypaInstallHook,
|
||||
wheelUnpackHook,
|
||||
cudaPackages,
|
||||
python,
|
||||
jaxlib,
|
||||
}:
|
||||
let
|
||||
inherit (jaxlib) version;
|
||||
|
||||
platforms = {
|
||||
x86_64-linux = {
|
||||
name = "manylinux_2_27_x86_64";
|
||||
hash = "sha256-TKiiYaVIWtuFL96VTkuRRfVCwNivRwSojgX67jVBryE=";
|
||||
};
|
||||
aarch64-linux = {
|
||||
name = "manylinux_2_27_aarch64";
|
||||
hash = "sha256-gZx0TjXd8Cckv0WWpijSvjhzaah82mBXUTfuQ37NeXo=";
|
||||
};
|
||||
};
|
||||
currentPlatform = platforms.${stdenv.hostPlatform.system};
|
||||
|
||||
cudaLibPath = lib.makeLibraryPath (
|
||||
with cudaPackages;
|
||||
[
|
||||
(lib.getLib libcublas) # libcublas.so
|
||||
(lib.getLib cuda_cupti) # libcupti.so
|
||||
(lib.getLib cuda_cudart) # libcudart.so
|
||||
(lib.getLib cudnn) # libcudnn.so
|
||||
(lib.getLib libcufft) # libcufft.so
|
||||
(lib.getLib libcusolver) # libcusolver.so
|
||||
(lib.getLib libcusparse) # libcusparse.so
|
||||
(lib.getLib nccl) # libnccl.so
|
||||
(lib.getLib libnvjitlink) # libnvJitLink.so
|
||||
(lib.getLib addDriverRunpath.driverLink) # libcuda.so
|
||||
]
|
||||
);
|
||||
|
||||
in
|
||||
buildPythonPackage (finalAttrs: {
|
||||
pname = "jax-cuda13-pjrt";
|
||||
inherit version;
|
||||
pyproject = false;
|
||||
__structuredAttrs = true;
|
||||
|
||||
src = fetchPypi {
|
||||
pname = "jax_cuda13_pjrt";
|
||||
inherit version;
|
||||
format = "wheel";
|
||||
python = "py3";
|
||||
dist = "py3";
|
||||
platform = currentPlatform.name;
|
||||
inherit (currentPlatform) hash;
|
||||
};
|
||||
|
||||
nativeBuildInputs = [
|
||||
autoPatchelfHook
|
||||
pypaInstallHook
|
||||
wheelUnpackHook
|
||||
];
|
||||
|
||||
# jax-cuda13-pjrt looks for ptxas, nvlink and nvvm at runtime, eg when running `jax.random.PRNGKey(0)`.
|
||||
# Linking into $out is the least bad solution. See
|
||||
# * https://github.com/NixOS/nixpkgs/pull/164176#discussion_r828801621
|
||||
# * https://github.com/NixOS/nixpkgs/pull/288829#discussion_r1493852211
|
||||
# for more info.
|
||||
postInstall = ''
|
||||
export OUTPATH="$out/${python.sitePackages}/jax_plugins/nvidia/cu13"
|
||||
export BINPATH="$OUTPATH/bin"
|
||||
mkdir -p $BINPATH
|
||||
ln -s ${lib.getExe' cudaPackages.cuda_nvcc "ptxas"} $BINPATH/ptxas
|
||||
ln -s ${lib.getExe' cudaPackages.cuda_nvcc "nvlink"} $BINPATH/nvlink
|
||||
ln -s ${cudaPackages.cuda_nvcc}/nvvm $OUTPATH/nvvm
|
||||
'';
|
||||
|
||||
# jax-cuda13-pjrt contains shared libraries that open other shared libraries via dlopen
|
||||
# and these implicit dependencies are not recognized by ldd or
|
||||
# autoPatchelfHook. That means we need to sneak them into rpath. This step
|
||||
# must be done after autoPatchelfHook and the automatic stripping of
|
||||
# artifacts. autoPatchelfHook runs in postFixup and auto-stripping runs in the
|
||||
# patchPhase.
|
||||
preInstallCheck = ''
|
||||
patchelf --add-rpath "${cudaLibPath}" $out/${python.sitePackages}/jax_plugins/xla_cuda13/xla_cuda_plugin.so
|
||||
'';
|
||||
|
||||
# FIXME: there are no tests, but we need to run preInstallCheck above
|
||||
doCheck = true;
|
||||
|
||||
pythonImportsCheck = [ "jax_plugins" ];
|
||||
|
||||
passthru = {
|
||||
inherit cudaLibPath;
|
||||
};
|
||||
|
||||
meta = {
|
||||
description = "JAX XLA PJRT Plugin for NVIDIA GPUs";
|
||||
homepage = "https://github.com/jax-ml/jax/tree/main/jax_plugins/cuda";
|
||||
sourceProvenance = [ lib.sourceTypes.binaryNativeCode ];
|
||||
license = lib.licenses.asl20;
|
||||
teams = [ lib.teams.cuda ];
|
||||
maintainers = with lib.maintainers; [ GaetanLepage ];
|
||||
platforms = lib.attrNames platforms;
|
||||
problems =
|
||||
lib.optionalAttrs (cudaPackages.cudaMajorVersion != "13") {
|
||||
unsupported-cuda-version = {
|
||||
message = ''
|
||||
Incompatible cudaPackages version.
|
||||
- Expected: 13
|
||||
- Got: ${cudaPackages.cudaMajorVersion}
|
||||
'';
|
||||
kind = "broken";
|
||||
};
|
||||
}
|
||||
// lib.optionalAttrs (lib.versionAtLeast cudaPackages.cudnn.version "10.0") {
|
||||
unsupported-cudnn-version = {
|
||||
message = ''
|
||||
cudaPackages.cudnn is too new (${cudaPackages.cudnn.version}).
|
||||
|
||||
See CUDA compatibility matrix
|
||||
https://docs.jax.dev/en/latest/installation.html#pip-installation-nvidia-gpu-cuda-installed-locally-harder
|
||||
'';
|
||||
kind = "broken";
|
||||
};
|
||||
};
|
||||
};
|
||||
})
|
||||
@@ -20,4 +20,5 @@ done
|
||||
|
||||
for arch in "x86_64-linux" "aarch64-linux"; do
|
||||
prefetch "312" "$arch" "jax-cuda12-pjrt"
|
||||
prefetch "312" "$arch" "jax-cuda13-pjrt"
|
||||
done
|
||||
|
||||
@@ -8692,6 +8692,8 @@ self: super: with self; {
|
||||
|
||||
jax-cuda12-plugin = callPackage ../development/python-modules/jax-cuda12-plugin { };
|
||||
|
||||
jax-cuda13-pjrt = callPackage ../development/python-modules/jax-cuda13-pjrt { };
|
||||
|
||||
jax-jumpy = callPackage ../development/python-modules/jax-jumpy { };
|
||||
|
||||
jax-tap = callPackage ../development/python-modules/jax-tap { };
|
||||
|
||||
Reference in New Issue
Block a user