include cudnn as python dep

This commit is contained in:
Awni Hannun
2025-07-18 06:54:41 -07:00
committed by Cheng
parent 180ec0d3a5
commit 75bcb46069
2 changed files with 2 additions and 1 deletions

View File

@@ -16,7 +16,7 @@ rm "${repaired_wheel}"
mlx_so="mlx/lib/libmlx.so"
rpath=$(patchelf --print-rpath "${mlx_so}")
base="\$ORIGIN/../../nvidia"
rpath=$rpath:${base}/cublas/lib:${base}/cuda_nvrtc/lib
rpath=$rpath:${base}/cublas/lib:${base}/cuda_nvrtc/lib:${base}/cudnn/lib
patchelf --force-rpath --set-rpath "$rpath" "$mlx_so"
python ../python/scripts/repair_record.py ${mlx_so}

View File

@@ -289,6 +289,7 @@ if __name__ == "__main__":
install_requires += [
"nvidia-cublas-cu12==12.9.*",
"nvidia-cuda-nvrtc-cu12==12.9.*",
"nvidia-cudnn-cu12==12.9.*",
]
else:
name = "mlx-cpu"