load("@bazel_lib//lib:copy_to_directory.bzl", "copy_to_directory") load("@zml//bazel:patchelf.bzl", "patchelf") load("@rules_cc//cc:cc_library.bzl", "cc_library") load("@zml//platforms/cuda:cuda.bzl", "CUDA_COMPAT_FILES") patchelf( name = "libpjrt_cuda_so", src = "libpjrt_cuda.so", add_needed = [ "libzmlxcuda.so.0", ], rename_dynamic_symbols = { "dlopen": "zmlxcuda_dlopen", }, replace_needed = { "nvshmem_transport_ibrc.so.3": "nvshmem_transport_ibrc.so.4", "libnvrtc-builtins.so.13.0": "libnvrtc-builtins.so.13.1", }, set_rpath = "$ORIGIN", ) copy_to_directory( name = "sandbox", srcs = [ ":libpjrt_cuda_so", "@zml//platforms/cuda:compat_probe", "@zml//platforms/cuda:zmlxcuda", ] + select({ "@llvm//platforms/config:linux_x86_64": [ "@cuda_cudart_linux_x86_64//:cuda_cudart", "@cuda_cupti_linux_x86_64//:cuda_cupti", "@cuda_compat_linux_x86_64//:cuda_compat", "@cuda_nvcc_linux_x86_64//:cuda_nvcc", "@libnvvm_linux_x86_64//:libnvvm", "@cuda_nvrtc_linux_x86_64//:cuda_nvrtc", "@cuda_nvtx_linux_x86_64//:cuda_nvtx", "@cudnn_linux_x86_64//:cudnn", "@libnvshmem_linux_x86_64//:libnvshmem", "@libcublas_linux_x86_64//:libcublas", "@libcufft_linux_x86_64//:libcufft", "@libcusolver_linux_x86_64//:libcusolver", "@libcusparse_linux_x86_64//:libcusparse", "@libnvjitlink_linux_x86_64//:libnvjitlink", "@nccl_linux_amd64//:nccl", "@zlib1g_linux_amd64//:zlib1g", ], "@llvm//platforms/config:linux_aarch64": [ "@cuda_cudart_linux_sbsa//:cuda_cudart", "@cuda_cupti_linux_sbsa//:cuda_cupti", "@cuda_compat_linux_sbsa//:cuda_compat", "@cuda_nvcc_linux_sbsa//:cuda_nvcc", "@libnvvm_linux_sbsa//:libnvvm", "@cuda_nvrtc_linux_sbsa//:cuda_nvrtc", "@cuda_nvtx_linux_sbsa//:cuda_nvtx", "@cudnn_linux_sbsa//:cudnn", "@libnvshmem_linux_sbsa//:libnvshmem", "@libcublas_linux_sbsa//:libcublas", "@libcufft_linux_sbsa//:libcufft", "@libcusolver_linux_sbsa//:libcusolver", "@libcusparse_linux_sbsa//:libcusparse", "@libnvjitlink_linux_sbsa//:libnvjitlink", "@nccl_linux_arm64//:nccl", "@zlib1g_linux_arm64//:zlib1g", ], }), replace_prefixes = { "compat": "lib/compat", "nvidia/nccl/lib": "lib", "nvvm/lib64": "lib", "libpjrt_cuda_so": "lib", "lib/x86_64-linux-gnu": "lib", "lib/aarch64-linux-gnu": "lib", "platforms/cuda/compat_probe": "bin/compat_probe", "platforms/cuda/libzmlxcuda": "lib/libzmlxcuda", } | { "{}.patchelf".format(file): "lib/compat" for file in CUDA_COMPAT_FILES }, add_directory_to_runfiles = False, include_external_repositories = ["**"], ) cc_library( name = "libpjrt_cuda", data = [":sandbox"], deps = select({ "@llvm//platforms/config:linux_x86_64": [ "@cuda_cudart_linux_x86_64//:cuda", ], "@llvm//platforms/config:linux_aarch64": [ "@cuda_cudart_linux_sbsa//:cuda", ], }), linkopts = [ # Defer function call resolution until the function is called # (lazy loading) rather than at load time. # # This is required because we want to let downstream use weak CUDA symbols. # # We force it here because -z,now (which resolve all symbols at load time), # is the default in most bazel CC toolchains as well as in certain linkers. "-Wl,-z,lazy", ], visibility = ["@zml//platforms/cuda:__subpackages__"], )