load("@rules_cc//cc:cc_library.bzl", "cc_library") load("@rules_zig//zig:defs.bzl", "zig_library", "zig_shared_library", "zig_binary"); zig_binary( name = "compat_probe", main = "compat_probe.zig", visibility = ["//visibility:public"], ) zig_shared_library( name = "zmlxcuda", main = "zmlxcuda.zig", shared_lib_name = "libzmlxcuda.so.0", visibility = ["//visibility:public"], # Use Clang's compiler-rt, but disable stack checking # to avoid requiring on the _zig_probe_stack symbol. zigopts = ["-fno-stack-check"], deps = ["//stdx"], ) cc_library( name = "empty", tags = ["manual"], defines = ["ZML_RUNTIME_CUDA_DISABLED"], ) cc_library( name = "libpjrt_cuda_repo", deps = select({ "@llvm//platforms/config:linux_x86_64": ["@libpjrt_cuda_linux_amd64//:libpjrt_cuda"], "@llvm//platforms/config:linux_aarch64": ["@libpjrt_cuda_linux_arm64//:libpjrt_cuda"], }), ) cc_library( name = "cuda_nvtx_headers", visibility = ["//visibility:public"], deps = select({ "@llvm//platforms/config:linux_x86_64": ["@cuda_nvtx_linux_x86_64//:headers"], "@llvm//platforms/config:linux_aarch64": ["@cuda_nvtx_linux_sbsa//:headers"], }), ) cc_library( name = "libpjrt_cuda", tags = ["manual"], hdrs = ["libpjrt_cuda.h"], defines = ["ZML_RUNTIME_CUDA"], deps = [":libpjrt_cuda_repo"], ) zig_library( name = "cuda", tags = ["manual"], srcs = ["compat_probe.zig"], import_name = "platforms/cuda", main = "cuda.zig", visibility = ["//visibility:public"], deps = [ "//pjrt", ] + select({ "//platforms:cuda.enabled": [ ":libpjrt_cuda", "//bazel", "//stdx", "@rules_zig//zig/runfiles", ], "//conditions:default": [":empty"], }), )