diff --git a/build/build.py b/build/build.py index 40e02a100d98..dab92162e301 100755 --- a/build/build.py +++ b/build/build.py @@ -72,16 +72,14 @@ "jax_source_package": "//:jax_source_package", "jaxlib": "//jaxlib/tools:jaxlib_wheel", "jaxlib_editable": "//jaxlib/tools:jaxlib_wheel_editable", - "jax-cuda-plugin": "//jaxlib/tools:jax_cuda_plugin_wheel", - "jax-cuda-plugin_editable": "//jaxlib/tools:jax_cuda_plugin_wheel_editable", - "jax-cuda-pjrt": "//jaxlib/tools:jax_cuda_pjrt_wheel", - "jax-cuda-pjrt_editable": "//jaxlib/tools:jax_cuda_pjrt_wheel_editable", + "jax-cuda-plugin": "//jaxlib/tools:jax_cuda{cuda_major_version}_plugin_wheel", + "jax-cuda-plugin_editable": "//jaxlib/tools:jax_cuda{cuda_major_version}_plugin_wheel_editable", + "jax-cuda-pjrt": "//jaxlib/tools:jax_cuda{cuda_major_version}_pjrt_wheel", + "jax-cuda-pjrt_editable": "//jaxlib/tools:jax_cuda{cuda_major_version}_pjrt_wheel_editable", "jax-rocm-plugin": "//jaxlib/tools:jax_rocm_plugin_wheel", "jax-rocm-pjrt": "//jaxlib/tools:jax_rocm_pjrt_wheel", } -_JAX_CUDA_VERSION = "12" - def add_global_arguments(parser: argparse.ArgumentParser): """Adds all the global arguments that applies to all the CLI subcommands.""" parser.add_argument( @@ -642,6 +640,11 @@ async def main(): # https://peps.python.org/pep-0440/ wheel_git_hash = option.split("=")[-1].lstrip('0')[:9] + if args.cuda_version: + cuda_major_version = args.cuda_version.split(".")[0] + else: + cuda_major_version = args.cuda_major_version + with open(".jax_configure.bazelrc", "w") as f: jax_configure_options = utils.get_jax_configure_bazel_options(wheel_build_command_base.get_command_as_list(), args.use_new_wheel_build_rule) if not jax_configure_options: @@ -692,6 +695,7 @@ async def main(): build_target = wheel_build_targets[wheel + "_editable"] else: build_target = wheel_build_targets[wheel] + build_target = build_target.format(cuda_major_version=cuda_major_version) wheel_build_command.append(build_target) if args.use_new_wheel_build_rule and wheel == "jax" and not args.editable: wheel_build_command.append(wheel_build_targets["jax_source_package"]) @@ -709,10 +713,6 @@ async def main(): if "cuda" in wheel: wheel_build_command.append("--enable-cuda=True") - if args.cuda_version: - cuda_major_version = args.cuda_version.split(".")[0] - else: - cuda_major_version = args.cuda_major_version wheel_build_command.append(f"--platform_version={cuda_major_version}") if "rocm" in wheel: @@ -738,7 +738,7 @@ async def main(): else: bazel_dir = jaxlib_and_plugins_bazel_dir if "cuda" in wheel: - wheel_dir = wheel.replace("cuda", f"cuda{_JAX_CUDA_VERSION}").replace( + wheel_dir = wheel.replace("cuda", f"cuda{cuda_major_version}").replace( "-", "_" ) else: diff --git a/build/requirements.in b/build/requirements.in index a88c194f7b8e..dd80114978e8 100644 --- a/build/requirements.in +++ b/build/requirements.in @@ -20,7 +20,9 @@ jaxlib==0.6.2 # The with-cuda extra also includes NVIDIA's pip packages. jax-cuda12-plugin[with-cuda]==0.6.2 ; sys_platform == "linux" +jax-cuda13-plugin jax-cuda12-pjrt==0.6.2 ; sys_platform == "linux" +jax-cuda13-pjrt # TPU dependencies libtpu ; sys_platform == "linux" and platform_machine == "x86_64" @@ -28,6 +30,7 @@ libtpu ; sys_platform == "linux" and platform_machine == "x86_64" # For Mosaic GPU collectives nvidia-cuda-nvrtc-cu12>=12.1.55 ; sys_platform == "linux" nvidia-nvshmem-cu12>=3.2.5 ; sys_platform == "linux" +nvidia-nvshmem-cu13 # Platform-specific dependencies that are being ignored by pip-compile colorama>=0.4.4 diff --git a/build/requirements_lock_3_12.txt b/build/requirements_lock_3_12.txt index 743fbbba325f..9f8ac8db50ed 100644 --- a/build/requirements_lock_3_12.txt +++ b/build/requirements_lock_3_12.txt @@ -172,6 +172,12 @@ jax-cuda12-plugin[with-cuda]==0.6.2 ; sys_platform == "linux" \ --hash=sha256:ed5316ca1818db7ef53230ee0a41398d3a60942e361dfb857a952eb4d92fc8d7 \ --hash=sha256:febd099f970d350eb8fa5a2c9a2fb4b0ea7b3d6a89df1496663edfa7afe590e5 # via -r build/requirements.in +jax-cuda13-pjrt==0.0.1rc0 \ + --hash=sha256:60835b87c9e1e5e1109ef2e9f27db6c9404af3ce3c3315316fb3b113cb13a158 + # via -r build/requirements.in +jax-cuda13-plugin==0.0.1rc0 \ + --hash=sha256:5fa66ef5de34cffc838199d193bffc11701807e66831f914bdc1d786ed4dc26f + # via -r build/requirements.in jaxlib==0.6.2 \ --hash=sha256:11eae7e05bc5a79875da36324afb9eddd4baeaef2a0386caf6d4f3720b9aef28 \ --hash=sha256:153eaa51f778b60851720729d4f461a91edd9ba3932f6f3bc598d4413870038b \ @@ -499,6 +505,10 @@ nvidia-nvshmem-cu12==3.2.5 ; sys_platform == "linux" \ # via # -r build/requirements.in # jax-cuda12-plugin +nvidia-nvshmem-cu13==0.0.0a0 \ + --hash=sha256:84d265d7b97dae6ee74139f8f7e37fc65a63e4ebb7287b987a4dca0c0625673d \ + --hash=sha256:b6900e44e6be1e0e7be6059c5b7a397fb3cb84914784571ab7e20a35bb2b140d + # via -r build/requirements.in opt-einsum==3.3.0 \ --hash=sha256:2455e59e3947d3c275477df7f5205b30635e266fe6dc300e3d9f9646bfcea147 \ --hash=sha256:59f6475f77bbc37dcf7cd748519c0ec60722e91e63ca114e68821c0c54a46549 diff --git a/jax/_src/lib/__init__.py b/jax/_src/lib/__init__.py index fde926094e8b..8fb2b85460ef 100644 --- a/jax/_src/lib/__init__.py +++ b/jax/_src/lib/__init__.py @@ -17,10 +17,12 @@ from __future__ import annotations +import importlib import gc import os import pathlib import re +from types import ModuleType try: import jaxlib as jaxlib @@ -119,13 +121,16 @@ def _xla_gc_callback(*args): xla_client._xla.collect_garbage() gc.callbacks.append(_xla_gc_callback) -try: - import jaxlib.cuda._versions as cuda_versions # pytype: disable=import-error # noqa: F401 -except ImportError: +cuda_versions: ModuleType | None +for pkg_name in ['jax_cuda13_plugin', 'jax_cuda12_plugin', 'jaxlib.cuda']: try: - import jax_cuda12_plugin._versions as cuda_versions # pytype: disable=import-error # noqa: F401 + cuda_versions = importlib.import_module( + f'{pkg_name}._versions' + ) except ImportError: cuda_versions = None + else: + break import jaxlib.gpu_solver as gpu_solver # pytype: disable=import-error # noqa: F401 import jaxlib.gpu_sparse as gpu_sparse # pytype: disable=import-error # noqa: F401 diff --git a/jax/_src/lib/mosaic_gpu.py b/jax/_src/lib/mosaic_gpu.py index 494112093029..37c190a409c5 100644 --- a/jax/_src/lib/mosaic_gpu.py +++ b/jax/_src/lib/mosaic_gpu.py @@ -18,6 +18,9 @@ try: from jaxlib.mosaic.gpu import _mosaic_gpu_ext # pytype: disable=import-error except ImportError: - from jax_cuda12_plugin import _mosaic_gpu_ext # pytype: disable=import-error + try: + from jax_cuda12_plugin import _mosaic_gpu_ext # pytype: disable=import-error + except ImportError: + from jax_cuda13_plugin import _mosaic_gpu_ext # pytype: disable=import-error except ImportError as e: raise ModuleNotFoundError("Failed to import the Mosaic GPU bindings") from e diff --git a/jax/_src/numpy/array_constructors.py b/jax/_src/numpy/array_constructors.py index 73bbd7d09554..8f25adc57802 100644 --- a/jax/_src/numpy/array_constructors.py +++ b/jax/_src/numpy/array_constructors.py @@ -32,7 +32,7 @@ export = util.set_module('jax.numpy') -for pkg_name in ['jax_cuda12_plugin', 'jax.jaxlib.cuda']: +for pkg_name in ['jax_cuda13_plugin', 'jax_cuda12_plugin', 'jax.jaxlib.cuda']: try: cuda_plugin_extension = importlib.import_module( f'{pkg_name}.cuda_plugin_extension' diff --git a/jax/experimental/mosaic/gpu/core.py b/jax/experimental/mosaic/gpu/core.py index 4d4193082e5b..1a41bad9f057 100644 --- a/jax/experimental/mosaic/gpu/core.py +++ b/jax/experimental/mosaic/gpu/core.py @@ -884,7 +884,10 @@ def as_torch_gpu_kernel( # Get our hands on the compilation and unload functions try: - import jax_plugins.xla_cuda12 as cuda_plugin # pytype: disable=import-error + try: + import jax_plugins.xla_cuda13 as cuda_plugin # pytype: disable=import-error + except ImportError: + import jax_plugins.xla_cuda12 as cuda_plugin # pytype: disable=import-error except ImportError: raise RuntimeError("as_torch_gpu_kernel only works with recent jaxlib builds " "that use backend plugins") diff --git a/jax/experimental/pallas/ops/gpu/attention_mgpu.py b/jax/experimental/pallas/ops/gpu/attention_mgpu.py index 90b8eb702db4..bdf0bda2a725 100644 --- a/jax/experimental/pallas/ops/gpu/attention_mgpu.py +++ b/jax/experimental/pallas/ops/gpu/attention_mgpu.py @@ -63,6 +63,7 @@ def has_backward_blocks(self) -> bool: return self.block_q_dkv is not None def _attention_forward(q, k, v, config: TuningConfig, save_residuals: bool = False): + assert cuda_versions is not None cuda_runtime_version = cuda_versions.cuda_runtime_get_version() # TODO(pobudzey): Undo when we upgrade to cuda 12.9.1. if config.causal and cuda_runtime_version >= 12080 and cuda_runtime_version < 12091: diff --git a/jax_plugins/cuda/__init__.py b/jax_plugins/cuda/__init__.py index de296a7a9e81..ca7154e6d6f5 100644 --- a/jax_plugins/cuda/__init__.py +++ b/jax_plugins/cuda/__init__.py @@ -34,7 +34,7 @@ def _import_extensions(): # cuda_plugin_extension locates inside jaxlib. `jaxlib` is for testing without # preinstalled jax cuda plugin packages. - for pkg_name in ['jax_cuda12_plugin', 'jaxlib.cuda']: + for pkg_name in ['jax_cuda13_plugin', 'jax_cuda12_plugin', 'jaxlib.cuda']: try: cuda_plugin_extension = importlib.import_module( f'{pkg_name}.cuda_plugin_extension' @@ -124,16 +124,23 @@ def _load_nvidia_libraries(): them from LD_LIBRARY_PATH. By loading the libraries here, later lookups will find these copies.""" _load("cuda_runtime", ["libcudart.so.12"]) + _load("cu13", ["libcudart.so.13"]) # cuda_nvrtc isn't directly a dependency of JAX, but CUDNN appears to need it # and at least in CUDA 12.9 has RUNPATHs misconfigured to refer to # nvidia/nvrtc instead of nvidia/cuda_nvrtc. _load("cuda_nvrtc", ["libnvrtc.so.12"]) + _load("cu13", ["libnvrtc.so.13"]) _load("cublas", ["libcublas.so.12", "libcublasLt.so.12"]) + _load("cu13", ["libcublas.so.13", "libcublasLt.so.13"]) _load("nccl", ["libnccl.so.2"]) _load("cuda_cupti", ["libcupti.so.12"]) + _load("cu13", ["libcupti.so.13"]) _load("cusparse", ["libcusparse.so.12"]) + _load("cu13", ["libcusparse.so.12"]) _load("cusolver", ["libcusolver.so.11"]) + _load("cu13", ["libcusolver.so.12"]) _load("cufft", ["libcufft.so.11"]) + _load("cu13", ["libcufft.so.12"]) _load("nvshmem", ["libnvshmem_host.so.3"]) _load("cudnn", ["libcudnn.so.9"]) diff --git a/jax_plugins/cuda/plugin_setup.py b/jax_plugins/cuda/plugin_setup.py index baa20f2419fc..9d86adc1cac4 100644 --- a/jax_plugins/cuda/plugin_setup.py +++ b/jax_plugins/cuda/plugin_setup.py @@ -21,6 +21,7 @@ cuda_version = 0 # placeholder project_name = f"jax-cuda{cuda_version}-plugin" package_name = f"jax_cuda{cuda_version}_plugin" +cuda_whl_sfx = "-cu12" if cuda_version == 12 else "" def load_version_module(pkg_path): spec = importlib.util.spec_from_file_location( @@ -53,15 +54,15 @@ def has_ext_modules(self): install_requires=[f"jax-cuda{cuda_version}-pjrt=={__version__}"], extras_require={ 'with-cuda': [ - "nvidia-cublas-cu12>=12.1.3.1", - "nvidia-cuda-cupti-cu12>=12.1.105", - "nvidia-cuda-nvcc-cu12>=12.6.85", - "nvidia-cuda-runtime-cu12>=12.1.105", - "nvidia-cudnn-cu12>=9.8,<10.0", - "nvidia-cufft-cu12>=11.0.2.54", - "nvidia-cusolver-cu12>=11.4.5.107", - "nvidia-cusparse-cu12>=12.1.0.106", - "nvidia-nccl-cu12>=2.18.1", + f"nvidia-cublas{cuda_whl_sfx}>=12.1.3.1", + f"nvidia-cuda-cupti{cuda_whl_sfx}>=12.1.105", + f"nvidia-cuda-nvcc{cuda_whl_sfx}>=12.6.85", + f"nvidia-cuda-runtime{cuda_whl_sfx}>=12.1.105", + f"nvidia-cudnn-cu{cuda_version}>=9.8,<10.0", + f"nvidia-cufft{cuda_whl_sfx}>=11.0.2.54", + f"nvidia-cusolver{cuda_whl_sfx}>=11.4.5.107", + f"nvidia-cusparse{cuda_whl_sfx}>=12.1.0.106", + f"nvidia-nccl-cu{cuda_version}>=2.18.1", # nvjitlink is not a direct dependency of JAX, but it is a transitive # dependency via, for example, cuSOLVER. NVIDIA's cuSOLVER packages # do not have a version constraint on their dependencies, so the @@ -69,13 +70,13 @@ def has_ext_modules(self): # problems (https://github.com/jax-ml/jax/issues/18027#issuecomment-1756305196) # Until NVIDIA add version constraints, add a version constraint # here. - "nvidia-nvjitlink-cu12>=12.1.105", + f"nvidia-nvjitlink{cuda_whl_sfx}>=12.1.105", # nvrtc is a transitive and undeclared dep of cudnn. - "nvidia-cuda-nvrtc-cu12>=12.1.55", + f"nvidia-cuda-nvrtc{cuda_whl_sfx}>=12.1.55", # NVSHMEM is used by Mosaic GPU collectives and can be used by XLA to # speed up collectives too. - "nvidia-nvshmem-cu12>=3.2.5", - ], + f"nvidia-nvshmem-cu{cuda_version}>=3.2.5", + ] + (["nvidia-nvvm"] if cuda_version == 13 else []), }, url="https://github.com/jax-ml/jax", license="Apache-2.0", diff --git a/jaxlib/jax.bzl b/jaxlib/jax.bzl index 8e9500d2a1e0..03508b0d562f 100644 --- a/jaxlib/jax.bzl +++ b/jaxlib/jax.bzl @@ -26,6 +26,7 @@ load("@rules_python//python:defs.bzl", "py_library", "py_test") load("@xla//third_party/py:python_wheel.bzl", "collect_data_files", "transitive_py_deps") load("@xla//xla/tsl:tsl.bzl", "transitive_hdrs", _if_windows = "if_windows", _pybind_extension = "tsl_pybind_extension_opensource") load("@xla//xla/tsl/platform:build_config_root.bzl", _tf_cuda_tests_tags = "tf_cuda_tests_tags", _tf_exec_properties = "tf_exec_properties") +load("@cuda_cudart//:version.bzl", cuda_major_version = "VERSION") # Explicitly re-exports names to avoid "unused variable" warnings from .bzl # lint tools. @@ -188,7 +189,7 @@ def _gpu_test_deps(): "//jaxlib/rocm:gpu_only_test_deps", "//jax_plugins:gpu_plugin_only_test_deps", # TODO(ybaturina): Remove this once we can add NVSHMEM libraries in the dependencies. - "@pypi//nvidia_nvshmem_cu12", + "@pypi//nvidia_nvshmem_cu{cuda_major_version}".format(cuda_major_version=cuda_major_version), ], "//jax:config_build_jaxlib_false": [ "//jaxlib/tools:pypi_jax_cuda_plugin_with_cuda_deps", diff --git a/jaxlib/plugin_support.py b/jaxlib/plugin_support.py index ea24dc181be0..1d629d17e64a 100644 --- a/jaxlib/plugin_support.py +++ b/jaxlib/plugin_support.py @@ -21,9 +21,9 @@ from .version import __version__ as jaxlib_version -_PLUGIN_MODULE_NAME = { - "cuda": "jax_cuda12_plugin", - "rocm": "jax_rocm60_plugin", +_PLUGIN_MODULE_NAMES = { + "cuda": ["jax_cuda13_plugin", "jax_cuda12_plugin"], + "rocm": ["jax_rocm60_plugin"], } @@ -44,10 +44,10 @@ def import_from_plugin( The imported submodule, or None if the plugin is not installed or if the versions are incompatible. """ - if plugin_name not in _PLUGIN_MODULE_NAME: + if plugin_name not in _PLUGIN_MODULE_NAMES: raise ValueError(f"Unknown plugin: {plugin_name}") return maybe_import_plugin_submodule( - [f".{plugin_name}", _PLUGIN_MODULE_NAME[plugin_name]], + [f".{plugin_name}"] + _PLUGIN_MODULE_NAMES[plugin_name], submodule_name, check_version=check_version, ) diff --git a/jaxlib/tools/BUILD.bazel b/jaxlib/tools/BUILD.bazel index 30f5ede8c03c..7ee9c955bea8 100644 --- a/jaxlib/tools/BUILD.bazel +++ b/jaxlib/tools/BUILD.bazel @@ -36,6 +36,7 @@ load( "pytype_test", "wheel_sources", ) +load("@cuda_cudart//:version.bzl", cuda_major_version = "VERSION") licenses(["notice"]) # Apache 2 @@ -327,7 +328,7 @@ wheel_sources( ) jax_wheel( - name = "jax_cuda_plugin_wheel", + name = "jax_cuda12_plugin_wheel", enable_cuda = True, no_abi = False, # TODO(b/371217563) May use hermetic cuda version here. @@ -338,7 +339,18 @@ jax_wheel( ) jax_wheel( - name = "jax_cuda_plugin_wheel_editable", + name = "jax_cuda13_plugin_wheel", + enable_cuda = True, + no_abi = False, + # TODO(b/371217563) May use hermetic cuda version here. + platform_version = "13", + source_files = [":jax_plugin_sources"], + wheel_binary = ":build_gpu_kernels_wheel_tool", + wheel_name = "jax_cuda13_plugin", +) + +jax_wheel( + name = "jax_cuda12_plugin_wheel_editable", editable = True, enable_cuda = True, # TODO(b/371217563) May use hermetic cuda version here. @@ -348,6 +360,17 @@ jax_wheel( wheel_name = "jax_cuda12_plugin", ) +jax_wheel( + name = "jax_cuda13_plugin_wheel_editable", + editable = True, + enable_cuda = True, + # TODO(b/371217563) May use hermetic cuda version here. + platform_version = "13", + source_files = [":jax_plugin_sources"], + wheel_binary = ":build_gpu_kernels_wheel_tool", + wheel_name = "jax_cuda13_plugin", +) + jax_wheel( name = "jax_rocm_plugin_wheel", enable_rocm = True, @@ -414,7 +437,7 @@ wheel_sources( ) jax_wheel( - name = "jax_cuda_pjrt_wheel", + name = "jax_cuda12_pjrt_wheel", enable_cuda = True, no_abi = True, # TODO(b/371217563) May use hermetic cuda version here. @@ -425,7 +448,18 @@ jax_wheel( ) jax_wheel( - name = "jax_cuda_pjrt_wheel_editable", + name = "jax_cuda13_pjrt_wheel", + enable_cuda = True, + no_abi = True, + # TODO(b/371217563) May use hermetic cuda version here. + platform_version = "13", + source_files = [":jax_pjrt_sources"], + wheel_binary = ":build_gpu_plugin_wheel_tool", + wheel_name = "jax_cuda13_pjrt", +) + +jax_wheel( + name = "jax_cuda12_pjrt_wheel_editable", editable = True, enable_cuda = True, # TODO(b/371217563) May use hermetic cuda version here. @@ -435,6 +469,17 @@ jax_wheel( wheel_name = "jax_cuda12_pjrt", ) +jax_wheel( + name = "jax_cuda13_pjrt_wheel_editable", + editable = True, + enable_cuda = True, + # TODO(b/371217563) May use hermetic cuda version here. + platform_version = "13", + source_files = [":jax_pjrt_sources"], + wheel_binary = ":build_gpu_plugin_wheel_tool", + wheel_name = "jax_cuda13_pjrt", +) + jax_wheel( name = "jax_rocm_pjrt_wheel", enable_rocm = True, @@ -456,21 +501,22 @@ jax_wheel( ) # Py_import targets. +cuda_suffix = "_cu12" if cuda_major_version == "12" else "" filegroup( name = "nvidia_wheel_deps", srcs = [ - "@pypi_nvidia_cublas_cu12//:pkg", - "@pypi_nvidia_cuda_cupti_cu12//:pkg", - "@pypi_nvidia_cuda_nvcc_cu12//:pkg", - "@pypi_nvidia_cuda_nvrtc_cu12//:pkg", - "@pypi_nvidia_cuda_runtime_cu12//:pkg", - "@pypi_nvidia_cudnn_cu12//:pkg", - "@pypi_nvidia_cufft_cu12//:pkg", - "@pypi_nvidia_cusolver_cu12//:pkg", - "@pypi_nvidia_cusparse_cu12//:pkg", - "@pypi_nvidia_nccl_cu12//:pkg", - "@pypi_nvidia_nvjitlink_cu12//:pkg", - "@pypi_nvidia_nvshmem_cu12//:pkg", + "@pypi_nvidia_cublas{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cuda_cupti{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cuda_nvcc{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cuda_nvrtc{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cuda_runtime{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cudnn_cu{cuda}//:pkg".format(cuda=cuda_major_version), + "@pypi_nvidia_cufft{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cusolver{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_cusparse{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_nccl_cu{cuda}//:pkg".format(cuda=cuda_major_version), + "@pypi_nvidia_nvjitlink{cuda}//:pkg".format(cuda=cuda_suffix), + "@pypi_nvidia_nvshmem_cu{cuda}//:pkg".format(cuda=cuda_major_version), ], ) @@ -508,13 +554,13 @@ py_import( # The targets below are used for GPU tests with `--//jax:build_jaxlib=false`. py_import( name = "pypi_jax_cuda_plugin_with_cuda_deps", - wheel = "@pypi_jax_cuda12_plugin//:whl", + wheel = "@pypi_jax_cuda{cuda}_plugin//:whl".format(cuda=cuda_major_version), wheel_deps = if_pypi_cuda_wheel_deps([":nvidia_wheel_deps"]), ) py_import( name = "pypi_jax_cuda_pjrt_with_cuda_deps", - wheel = "@pypi_jax_cuda12_pjrt//:whl", + wheel = "@pypi_jax_cuda{cuda}_pjrt//:whl".format(cuda=cuda_major_version), wheel_deps = if_pypi_cuda_wheel_deps([":nvidia_wheel_deps"]), ) @@ -544,7 +590,7 @@ verify_manylinux_compliance_test( test_tags = [ "manual", ], - wheel = ":jax_cuda_plugin_wheel", + wheel = ":jax_cuda{cuda}_plugin_wheel".format(cuda=cuda_major_version), x86_64_compliance_tag = X86_64_MANYLINUX_TAG, ) @@ -555,7 +601,7 @@ verify_manylinux_compliance_test( test_tags = [ "manual", ], - wheel = ":jax_cuda_pjrt_wheel", + wheel = ":jax_cuda{cuda}_pjrt_wheel".format(cuda=cuda_major_version), x86_64_compliance_tag = X86_64_MANYLINUX_TAG, ) @@ -578,10 +624,10 @@ pytype_test( name = "jax_cuda_plugin_wheel_size_test", srcs = [":wheel_size_test.py"], args = [ - "--wheel-path=$(location :jax_cuda_plugin_wheel)", + "--wheel-path=$(location :jax_cuda{cuda}_plugin_wheel)".format(cuda=cuda_major_version), "--max-size-mib=20", ], - data = [":jax_cuda_plugin_wheel"], + data = [":jax_cuda{cuda}_plugin_wheel".format(cuda=cuda_major_version)], main = "wheel_size_test.py", tags = [ "manual", @@ -593,10 +639,10 @@ pytype_test( name = "jax_cuda_pjrt_wheel_size_test", srcs = [":wheel_size_test.py"], args = [ - "--wheel-path=$(location :jax_cuda_pjrt_wheel)", + "--wheel-path=$(location :jax_cuda{cuda}_pjrt_wheel)".format(cuda=cuda_major_version), "--max-size-mib=120", ], - data = [":jax_cuda_pjrt_wheel"], + data = [":jax_cuda{cuda}_pjrt_wheel".format(cuda=cuda_major_version)], main = "wheel_size_test.py", tags = [ "manual", diff --git a/pyproject.toml b/pyproject.toml index ff34488124e7..4d534e42584e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,6 +23,7 @@ module = [ "jax.experimental.jax2tf.tests.back_compat_testdata", "jax.experimental.jax2tf.tests.flax_models", "jax_cuda12_plugin.*", + "jax_cuda13_plugin.*", "jaxlib.cpu_feature_guard", "jaxlib.cuda.*", "jaxlib.mlir.*", diff --git a/setup.py b/setup.py index a8bcdee95091..46ce0e5e4928 100644 --- a/setup.py +++ b/setup.py @@ -96,6 +96,11 @@ def load_version_module(pkg_path): f"jax-cuda12-plugin[with-cuda]>={_current_jaxlib_version},<={_jax_version}", ], + 'cuda13': [ + f"jaxlib>={_current_jaxlib_version},<={_jax_version}", + f"jax-cuda13-plugin[with-cuda]>={_current_jaxlib_version},<={_jax_version}", + ], + # Target that does not depend on the CUDA pip wheels, for those who want # to use a preinstalled CUDA. 'cuda12-local': [ @@ -103,6 +108,11 @@ def load_version_module(pkg_path): f"jax-cuda12-plugin>={_current_jaxlib_version},<={_jax_version}", ], + 'cuda13-local': [ + f"jaxlib>={_current_jaxlib_version},<={_jax_version}", + f"jax-cuda13-plugin>={_current_jaxlib_version},<={_jax_version}", + ], + # ROCm support for ROCm 6.0 and above. 'rocm': [ f"jaxlib>={_current_jaxlib_version},<={_jax_version}",