@@ -39,6 +39,7 @@ from cuda.core._utils.cuda_utils import (
3939 driver,
4040 is_sequence,
4141)
42+ from cuda.core._utils.version import driver_version
4243from cuda.core.typing import CompilerBackendType, ObjectCodeFormatType
4344
4445if TYPE_CHECKING:
@@ -52,7 +53,6 @@ _keep_driver_in_stub: "cuda.bindings.driver.CUlinkState"
5253_keep_nvjitlink_in_stub: " cuda.bindings.nvjitlink.nvJitLinkHandle"
5354
5455ctypedef const char * const_char_ptr
55- ctypedef void * void_ptr
5656
5757__all__ = [" Linker" , " LinkerOptions" ]
5858
@@ -155,10 +155,21 @@ cdef class Linker:
155155
156156 def close(self ) -> None:
157157 """Destroy this linker."""
158+ cdef vector[cydriver.CUjit_option] empty_keys
159+ cdef vector[void*] empty_values
158160 if self._use_nvjitlink:
159161 self._nvjitlink_handle.reset()
160162 else:
163+ if self._drv_log_bufs is not None:
164+ if self._info_log is None:
165+ self._info_log = self .get_info_log()
166+ if self._error_log is None:
167+ self._error_log = self .get_error_log()
168+ # Destroy the CUlinkState before releasing storage referenced by it.
161169 self._culink_handle.reset()
170+ self._drv_jit_keys.swap(empty_keys )
171+ self._drv_jit_values.swap(empty_values )
172+ self._drv_log_bufs = None
162173
163174 @property
164175 def handle(self ) -> LinkerHandleT:
@@ -207,7 +218,8 @@ class LinkerOptions:
207218
208219 Since the linker may choose to use nvJitLink or the driver APIs as the linking backend ,
209220 not all options are applicable. When the system's installed nvJitLink is too old (<12.3),
210- or not installed , the driver APIs (cuLink ) will be used instead.
221+ not installed , or older than the CUDA driver major version , the driver APIs (cuLink )
222+ will be used instead.
211223
212224 Attributes
213225 ----------
@@ -473,8 +485,8 @@ cdef inline int Linker_init(Linker self, tuple object_codes, object options) exc
473485 cdef cydriver.CUlinkState c_raw_culink
474486 cdef Py_ssize_t c_num_opts , i
475487 cdef vector[const_char_ptr] c_str_opts
476- cdef vector[ cydriver.CUjit_option] c_jit_keys
477- cdef vector[void_ptr] c_jit_values
488+ cdef cydriver.CUjit_option* c_drv_jit_keys_ptr
489+ cdef void** c_drv_jit_values_ptr
478490
479491 self._options = options = check_or_create_options(LinkerOptions, options, " Linker options" )
480492
@@ -496,19 +508,24 @@ cdef inline int Linker_init(Linker self, tuple object_codes, object options) exc
496508 # the driver writes into via raw pointers during linking operations.
497509 self ._drv_log_bufs = formatted_options
498510 c_num_opts = len (option_keys)
499- c_jit_keys .resize(c_num_opts)
500- c_jit_values .resize(c_num_opts)
511+ self ._drv_jit_keys .resize(c_num_opts)
512+ self ._drv_jit_values .resize(c_num_opts)
501513 for i in range (c_num_opts):
502- c_jit_keys [i] = < cydriver.CUjit_option>< int > option_keys[i]
514+ self ._drv_jit_keys [i] = < cydriver.CUjit_option>< int > option_keys[i]
503515 val = formatted_options[i]
504516 if isinstance (val, bytearray):
505- c_jit_values [i] = < void * > PyByteArray_AS_STRING(val)
517+ self ._drv_jit_values [i] = < void * > PyByteArray_AS_STRING(val)
506518 else :
507- c_jit_values[i] = < void * >< intptr_t> int (val)
519+ self ._drv_jit_values[i] = < void * >< intptr_t> int (val)
520+ c_drv_jit_keys_ptr = self ._drv_jit_keys.data()
521+ c_drv_jit_values_ptr = self ._drv_jit_values.data()
508522 try :
509523 with nogil:
510524 HANDLE_RETURN(cydriver.cuLinkCreate(
511- < unsigned int > c_num_opts, c_jit_keys.data(), c_jit_values.data(), & c_raw_culink))
525+ < unsigned int > c_num_opts,
526+ c_drv_jit_keys_ptr,
527+ c_drv_jit_values_ptr,
528+ & c_raw_culink))
512529 except CUDAError as e:
513530 Linker_annotate_error_log(self , e)
514531 raise
@@ -622,11 +639,10 @@ cdef inline object Linker_link(Linker self, str target_type):
622639 raise
623640 code = (< char * > c_cubin_out)[:c_output_size]
624641
625- # Linking is complete; cache the decoded log strings and release
626- # the driver's raw bytearray buffers (no longer written to ).
642+ # Linking is complete; cache the decoded logs. cuLinkDestroy may still
643+ # dereference the raw log-buffer pointers, so retain them until close( ).
627644 self ._info_log = self .get_info_log()
628645 self ._error_log = self .get_error_log()
629- self ._drv_log_bufs = None
630646
631647 return ObjectCode._init(bytes(code), target_type, name = self ._options.name)
632648
@@ -680,12 +696,22 @@ def _decide_nvjitlink_or_driver() -> bool:
680696 from cuda.bindings._internal import nvjitlink
681697
682698 if _nvjitlink_has_version_symbol(nvjitlink ):
683- _use_nvjitlink_backend = True
684- return False # Use nvjitlink
685- warn_txt = (
686- f" {'nvJitLink*.dll' if sys.platform == 'win32' else 'libnvJitLink.so*'} is too old (<12.3)."
687- f" Therefore cuda.bindings.nvjitlink is not usable and {warn_txt_common} nvJitLink."
688- )
699+ nvjitlink_version = nvjitlink_module.version()
700+ driver_major = driver_version()[0 ]
701+ if driver_major <= nvjitlink_version[0 ]:
702+ _use_nvjitlink_backend = True
703+ return False # Use nvjitlink
704+
705+ warn_txt = (
706+ f" CUDA driver major version {driver_major} is newer than "
707+ f" nvJitLink major version {nvjitlink_version[0]}; therefore "
708+ f" {warn_txt_common} nvJitLink."
709+ )
710+ else :
711+ warn_txt = (
712+ f" {'nvJitLink*.dll' if sys.platform == 'win32' else 'libnvJitLink.so*'} is too old (<12.3)."
713+ f" Therefore cuda.bindings.nvjitlink is not usable and {warn_txt_common} nvJitLink."
714+ )
689715
690716 warn(warn_txt, stacklevel = 2 , category = RuntimeWarning )
691717 _use_nvjitlink_backend = False
0 commit comments