Coverage for cuda/core/_linker.pyx: 78.88%
374 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-19 01:12 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-19 01:12 +0000
1# SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2#
3# SPDX-License-Identifier: Apache-2.0
4"""Linking machinery for combining object codes.
6This module provides :class:`Linker` for linking one or more
7:class:`~cuda.core.ObjectCode` objects, with :class:`LinkerOptions` for
8configuration.
9"""
11from __future__ import annotations
13from cpython.bytearray cimport PyByteArray_AS_STRING
14from libc.stdint cimport intptr_t, uint32_t
15from libcpp.vector cimport vector
16from cuda.bindings cimport cydriver
17from cuda.bindings cimport cynvjitlink
19from ._resource_handles cimport (
20 as_cu,
21 as_py,
22 create_culink_handle,
23 create_nvjitlink_handle,
24)
25from cuda.core._utils.cuda_utils cimport HANDLE_RETURN, HANDLE_RETURN_NVJITLINK
27import sys
28from dataclasses import dataclass
29from typing import TYPE_CHECKING, Union
30from warnings import warn
32from cuda.pathfinder._optional_cuda_import import _optional_cuda_import
33from cuda.core._device import Device
34from cuda.core._module import ObjectCode
35from cuda.core._utils.clear_error_support import assert_type
36from cuda.core._utils.cuda_utils import (
37 CUDAError,
38 check_or_create_options,
39 driver,
40 is_sequence,
41)
42from cuda.core.typing import CompilerBackendType, ObjectCodeFormatType
44if TYPE_CHECKING:
45 import cuda.bindings.driver # no-cython-lint
46 import cuda.bindings.nvjitlink # no-cython-lint
48# Module-level annotations to ensure stubgen-pyx keeps the above imports in
49# the generated `.pyi` so that the LinkerHandleT forward references resolve.
50# These names are not assigned, so they only affect __annotations__.
51_keep_driver_in_stub: "cuda.bindings.driver.CUlinkState"
52_keep_nvjitlink_in_stub: "cuda.bindings.nvjitlink.nvJitLinkHandle"
54ctypedef const char* const_char_ptr
56__all__ = ["Linker", "LinkerOptions"]
58LinkerHandleT = Union["cuda.bindings.nvjitlink.nvJitLinkHandle", "cuda.bindings.driver.CUlinkState"]
61# =============================================================================
62# Principal class
63# =============================================================================
65cdef class Linker:
66 """Represent a linking machinery to link one or more object codes into
67 :class:`~cuda.core.ObjectCode`.
69 This object provides a unified interface to multiple underlying
70 linker libraries (such as nvJitLink or cuLink* from the CUDA driver).
72 Parameters
73 ----------
74 object_codes : :class:`~cuda.core.ObjectCode`
75 One or more ObjectCode objects to be linked.
76 options : :class:`LinkerOptions`, optional
77 Options for the linker. If not provided, default options will be used.
78 """
80 def __init__(self, *object_codes: ObjectCode, options: LinkerOptions | None = None):
81 Linker_init(self, object_codes, options) 1(QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
83 def link(self, target_type: ObjectCodeFormatType | str) -> ObjectCode:
84 """Link the provided object codes into a single output of the specified target type.
86 Parameters
87 ----------
88 target_type : ObjectCodeFormatType | str
89 The type of the target output. Must be either "cubin" or "ptx".
91 Returns
92 -------
93 :class:`~cuda.core.ObjectCode`
94 The linked object code of the specified target type.
96 .. note::
98 Ensure that input object codes were compiled with appropriate
99 flags for linking (e.g., relocatable device code enabled).
100 """
101 return Linker_link(self, str(target_type)) 1QuOqzArBsCmknDEFGHIvowxpyJKSLltMNcdefghiabj
103 def get_error_log(self) -> str:
104 """Get the error log generated by the linker.
106 Returns
107 -------
108 str
109 The error log.
110 """
111 # After link(), the decoded log is cached here.
112 if self._error_log is not None: 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
113 return self._error_log 1t
114 cdef cynvjitlink.nvJitLinkHandle c_h
115 cdef size_t c_log_size = 0 1PQuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
116 cdef char* c_log_ptr
117 if self._use_nvjitlink: 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
118 c_h = as_cu(self._nvjitlink_handle) 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
119 HANDLE_RETURN_NVJITLINK(c_h, cynvjitlink.nvJitLinkGetErrorLogSize(c_h, &c_log_size)) 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
120 log = bytearray(c_log_size) 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
121 if c_log_size > 0: 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
122 c_log_ptr = <char*>(<bytearray>log) 1Q
123 HANDLE_RETURN_NVJITLINK(c_h, cynvjitlink.nvJitLinkGetErrorLog(c_h, c_log_ptr)) 1Q
124 return log.decode("utf-8", errors="backslashreplace") 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
125 else:
126 return (<bytearray>self._drv_log_bufs[2]).decode(
127 "utf-8", errors="backslashreplace").rstrip('\x00')
129 def get_info_log(self) -> str:
130 """Get the info log generated by the linker.
132 Returns
133 -------
134 str
135 The info log.
136 """
137 # After link(), the decoded log is cached here.
138 if self._info_log is not None: 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
139 return self._info_log 1ut
140 cdef cynvjitlink.nvJitLinkHandle c_h
141 cdef size_t c_log_size = 0 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
142 cdef char* c_log_ptr
143 if self._use_nvjitlink: 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
144 c_h = as_cu(self._nvjitlink_handle) 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
145 HANDLE_RETURN_NVJITLINK(c_h, cynvjitlink.nvJitLinkGetInfoLogSize(c_h, &c_log_size)) 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
146 log = bytearray(c_log_size) 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
147 if c_log_size > 0: 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
148 c_log_ptr = <char*>(<bytearray>log) 1qrsmklab
149 HANDLE_RETURN_NVJITLINK(c_h, cynvjitlink.nvJitLinkGetInfoLog(c_h, c_log_ptr)) 1qrsmklab
150 return log.decode("utf-8", errors="backslashreplace") 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
151 else:
152 return (<bytearray>self._drv_log_bufs[0]).decode(
153 "utf-8", errors="backslashreplace").rstrip('\x00')
155 def close(self) -> None:
156 """Destroy this linker."""
157 cdef vector[cydriver.CUjit_option] empty_keys
158 cdef vector[void*] empty_values
159 if self._use_nvjitlink: 1Rcdefghiabj
160 self._nvjitlink_handle.reset() 1Rcdefghiabj
161 else:
162 if self._drv_log_bufs is not None:
163 if self._info_log is None:
164 self._info_log = self.get_info_log()
165 if self._error_log is None:
166 self._error_log = self.get_error_log()
167 # Destroy the CUlinkState before releasing storage referenced by it.
168 self._culink_handle.reset()
169 self._drv_jit_keys.swap(empty_keys)
170 self._drv_jit_values.swap(empty_values)
171 self._drv_log_bufs = None
173 @property
174 def handle(self) -> LinkerHandleT:
175 """Return the underlying handle object.
177 .. note::
179 The type of the returned object depends on the backend.
181 .. caution::
183 This handle is a Python object. To get the memory address of the underlying C
184 handle, call ``int(Linker.handle)``.
185 """
186 if self._use_nvjitlink: 1TR
187 return as_py(self._nvjitlink_handle) 1TR
188 else:
189 return as_py(self._culink_handle)
191 @classmethod
192 def which_backend(cls) -> CompilerBackendType:
193 """Return which linking backend will be used.
195 Returns :attr:`~CompilerBackendType.NVJITLINK` when the nvJitLink
196 library is available and meets the minimum version requirement,
197 otherwise :attr:`~CompilerBackendType.DRIVER`.
199 .. note::
201 Prefer letting :class:`Linker` decide. Query ``which_backend()``
202 only when you need to dispatch based on input format (for
203 example: choose PTX vs. LTOIR before constructing a
204 ``Linker``). The returned value names an implementation
205 detail whose support matrix may shift across CTK releases.
206 """
207 return CompilerBackendType.DRIVER if _decide_nvjitlink_or_driver() else CompilerBackendType.NVJITLINK 2gb) * O q z A r B s C m k n D E F G H I v o w x p y J R c d e f g h i a b j
210# =============================================================================
211# Supporting classes
212# =============================================================================
214@dataclass
215class LinkerOptions:
216 """Customizable options for configuring :class:`Linker`.
218 Since the linker may choose to use nvJitLink or the driver APIs as the linking backend,
219 not all options are applicable. When the system's installed nvJitLink is too old (<12.3),
220 or not installed, the driver APIs (cuLink) will be used instead.
222 Attributes
223 ----------
224 name : str, optional
225 Name of the linker. If the linking succeeds, the name is passed down to the generated :class:`ObjectCode`.
226 arch : str, optional
227 Pass the SM architecture value, such as ``sm_<CC>`` (for generating CUBIN) or
228 ``compute_<CC>`` (for generating PTX). If not provided, the current device's architecture
229 will be used.
230 max_register_count : int, optional
231 Maximum register count.
232 time : bool, optional
233 Print timing information to the info log.
234 Default: False.
235 verbose : bool, optional
236 Print verbose messages to the info log.
237 Default: False.
238 link_time_optimization : bool, optional
239 Perform link time optimization.
240 Default: False.
241 ptx : bool, optional
242 Emit PTX after linking instead of CUBIN; only supported with ``link_time_optimization=True``.
243 Default: False.
244 optimization_level : int, optional
245 Set optimization level. Only 0 and 3 are accepted.
246 debug : bool, optional
247 Generate debug information.
248 Default: False.
249 lineinfo : bool, optional
250 Generate line information.
251 Default: False.
252 ftz : bool, optional
253 Flush denormal values to zero.
254 Default: False.
255 prec_div : bool, optional
256 Use precise division.
257 Default: True.
258 prec_sqrt : bool, optional
259 Use precise square root.
260 Default: True.
261 fma : bool, optional
262 Use fast multiply-add.
263 Default: True.
264 kernels_used : [str | tuple[str] | list[str]], optional
265 Pass a kernel or sequence of kernels that are used; any not in the list can be removed.
266 variables_used : [str | tuple[str] | list[str]], optional
267 Pass a variable or sequence of variables that are used; any not in the list can be removed.
268 optimize_unused_variables : bool, optional
269 Assume that if a variable is not referenced in device code, it can be removed.
270 Default: False.
271 ptxas_options : [str | tuple[str] | list[str]], optional
272 Pass options to PTXAS.
273 split_compile : int, optional
274 Split compilation maximum thread count. Use 0 to use all available processors. Value of 1 disables split
275 compilation (default).
276 Default: 1.
277 split_compile_extended : int, optional
278 A more aggressive form of split compilation available in LTO mode only. Accepts a maximum thread count value.
279 Use 0 to use all available processors. Value of 1 disables extended split compilation (default). Note: This
280 option can potentially impact performance of the compiled binary.
281 Default: 1.
282 no_cache : bool, optional
283 Do not cache the intermediate steps of nvJitLink.
284 Default: False.
285 """
287 name: str | None = "<default linker>"
288 arch: str | None = None
289 max_register_count: int | None = None
290 time: bool | None = None
291 verbose: bool | None = None
292 link_time_optimization: bool | None = None
293 ptx: bool | None = None
294 optimization_level: int | None = None
295 debug: bool | None = None
296 lineinfo: bool | None = None
297 ftz: bool | None = None
298 prec_div: bool | None = None
299 prec_sqrt: bool | None = None
300 fma: bool | None = None
301 kernels_used: str | tuple[str] | list[str] | None = None
302 variables_used: str | tuple[str] | list[str] | None = None
303 optimize_unused_variables: bool | None = None
304 ptxas_options: str | tuple[str] | list[str] | None = None
305 split_compile: int | None = None
306 split_compile_extended: int | None = None
307 no_cache: bool | None = None
308 numba_debug: bool | None = None
310 def __post_init__(self) -> None:
311 _lazy_init() 1P#QuTVKSLlt$%W7UXYZ012398654MNRcdefghiabj
312 self._name = self.name.encode() 1P#QuTVKSLlt$%W7UXYZ012398654MNRcdefghiabj
314 def _prepare_nvjitlink_options(self, as_bytes: bool = False) -> list[bytes] | list[str]:
315 options = [] 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
317 if self.arch is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
318 options.append(f"-arch={self.arch}") 1QuTqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
319 else:
320 options.append("-arch=sm_" + "".join(f"{i}" for i in Device().compute_capability)) 1O
321 if self.max_register_count is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
322 options.append(f"-maxrregcount={self.max_register_count}") 1zWc
323 if self.time is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
324 options.append("-time") 1sb
325 if self.verbose: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
326 options.append("-verbose") 1q
327 if self.link_time_optimization: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
328 options.append("-lto") 1l
329 if self.ptx: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
330 options.append("-ptx") 1Vl
331 if self.optimization_level is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
332 options.append(f"-O{self.optimization_level}") 1A
333 if self.debug: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
334 options.append("-g") 1rW7da
335 if self.lineinfo: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
336 options.append("-lineinfo") 1B7e
337 if self.ftz is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
338 options.append(f"-ftz={'true' if self.ftz else 'false'}") 1FWf
339 if self.prec_div is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
340 options.append(f"-prec-div={'true' if self.prec_div else 'false'}") 1Gg
341 if self.prec_sqrt is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
342 options.append(f"-prec-sqrt={'true' if self.prec_sqrt else 'false'}") 1Hh
343 if self.fma is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
344 options.append(f"-fma={'true' if self.fma else 'false'}") 1Ii
345 if self.kernels_used is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
346 if isinstance(self.kernels_used, str): 1vow
347 options.append(f"-kernels-used={self.kernels_used}") 1v
348 elif isinstance(self.kernels_used, list): 1ow
349 for kernel in self.kernels_used: 1o
350 options.append(f"-kernels-used={kernel}") 1o
351 if self.variables_used is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
352 if isinstance(self.variables_used, str): 1xpy
353 options.append(f"-variables-used={self.variables_used}") 1x
354 elif isinstance(self.variables_used, list): 1py
355 for variable in self.variables_used: 1p
356 options.append(f"-variables-used={variable}") 1p
357 if self.optimize_unused_variables is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
358 options.append("-optimize-unused-variables") 1C
359 if self.ptxas_options is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
360 if isinstance(self.ptxas_options, str): 1mkn
361 options.append(f"-Xptxas={self.ptxas_options}") 1m
362 elif is_sequence(self.ptxas_options): 1kn
363 for opt in self.ptxas_options: 1kn
364 options.append(f"-Xptxas={opt}") 1kn
365 if self.split_compile is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
366 options.append(f"-split-compile={self.split_compile}") 1Dj
367 if self.split_compile_extended is not None: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
368 options.append(f"-split-compile-extended={self.split_compile_extended}") 1E
369 if self.no_cache is True: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
370 options.append("-no-cache") 1J
372 if as_bytes: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltW7MNRcdefghiabj
373 return [o.encode() for o in options] 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltWMNRcdefghiabj
374 else:
375 return options 17
377 def _prepare_driver_options(self) -> tuple[list[object], list[object]]:
378 formatted_options = [] 1UXYZ012398654
379 option_keys = [] 1UXYZ012398654
381 # allocate a fixed-sized buffer for each info/error log
382 size = 4194304 1UXYZ012398654
383 formatted_options.extend((bytearray(size), size, bytearray(size), size)) 1UXYZ012398654
384 option_keys.extend( 1UXYZ012398654
385 (
386 _driver.CUjit_option.CU_JIT_INFO_LOG_BUFFER, 1UXYZ012398654
387 _driver.CUjit_option.CU_JIT_INFO_LOG_BUFFER_SIZE_BYTES, 1UXYZ012398654
388 _driver.CUjit_option.CU_JIT_ERROR_LOG_BUFFER, 1UXYZ012398654
389 _driver.CUjit_option.CU_JIT_ERROR_LOG_BUFFER_SIZE_BYTES, 1UXYZ012398654
390 )
391 )
393 if self.arch is not None: 1UXYZ012398654
394 arch = self.arch.split("_")[-1].upper() 1U
395 formatted_options.append(getattr(_driver.CUjit_target, f"CU_TARGET_COMPUTE_{arch}")) 1U
396 option_keys.append(_driver.CUjit_option.CU_JIT_TARGET) 1U
397 if self.max_register_count is not None: 1UXYZ012398654
398 formatted_options.append(self.max_register_count) 1U
399 option_keys.append(_driver.CUjit_option.CU_JIT_MAX_REGISTERS) 1U
400 if self.time is not None: 1UXYZ012398654
401 raise ValueError("time option is not supported by the driver API") 19
402 if self.verbose: 1UXYZ01238654
403 formatted_options.append(1) 1U
404 option_keys.append(_driver.CUjit_option.CU_JIT_LOG_VERBOSE) 1U
405 if self.link_time_optimization: 1UXYZ01238654
406 formatted_options.append(1) 1U
407 option_keys.append(_driver.CUjit_option.CU_JIT_LTO) 1U
408 if self.ptx: 1UXYZ01238654
409 raise ValueError("ptx option is not supported by the driver API") 18
410 if self.optimization_level is not None: 1UXYZ0123654
411 formatted_options.append(self.optimization_level) 1U
412 option_keys.append(_driver.CUjit_option.CU_JIT_OPTIMIZATION_LEVEL) 1U
413 if self.debug: 1UXYZ0123654
414 formatted_options.append(1) 1U
415 option_keys.append(_driver.CUjit_option.CU_JIT_GENERATE_DEBUG_INFO) 1U
416 if self.lineinfo: 1UXYZ0123654
417 formatted_options.append(1) 1U
418 option_keys.append(_driver.CUjit_option.CU_JIT_GENERATE_LINE_INFO) 1U
419 if self.ftz is not None: 1UXYZ0123654
420 warn("ftz option is deprecated in the driver API", DeprecationWarning, stacklevel=3) 1X
421 if self.prec_div is not None: 1UXYZ0123654
422 warn("prec_div option is deprecated in the driver API", DeprecationWarning, stacklevel=3) 1Y
423 if self.prec_sqrt is not None: 1UXYZ0123654
424 warn("prec_sqrt option is deprecated in the driver API", DeprecationWarning, stacklevel=3) 1Z
425 if self.fma is not None: 1UXYZ0123654
426 warn("fma options is deprecated in the driver API", DeprecationWarning, stacklevel=3) 10
427 if self.kernels_used is not None: 1UXYZ0123654
428 warn("kernels_used is deprecated in the driver API", DeprecationWarning, stacklevel=3) 11
429 if self.variables_used is not None: 1UXYZ0123654
430 warn("variables_used is deprecated in the driver API", DeprecationWarning, stacklevel=3) 12
431 if self.optimize_unused_variables is not None: 1UXYZ0123654
432 warn("optimize_unused_variables is deprecated in the driver API", DeprecationWarning, stacklevel=3) 13
433 if self.ptxas_options is not None: 1UXYZ0123654
434 raise ValueError("ptxas_options option is not supported by the driver API") 16
435 if self.split_compile is not None: 1UXYZ012354
436 raise ValueError("split_compile option is not supported by the driver API") 15
437 if self.split_compile_extended is not None: 1UXYZ01234
438 raise ValueError("split_compile_extended option is not supported by the driver API") 14
439 if self.no_cache is True: 1UXYZ0123
440 formatted_options.append(_driver.CUjit_cacheMode.CU_JIT_CACHE_OPTION_NONE) 1U
441 option_keys.append(_driver.CUjit_option.CU_JIT_CACHE_MODE) 1U
443 return formatted_options, option_keys 1UXYZ0123
445 def as_bytes(self, backend: str = "nvjitlink") -> list[bytes]:
446 """Convert linker options to bytes format for the nvjitlink backend.
448 Parameters
449 ----------
450 backend : str, optional
451 The linker backend. Only "nvjitlink" is supported. Default is "nvjitlink".
453 Returns
454 -------
455 list[bytes]
456 List of option strings encoded as bytes.
458 Raises
459 ------
460 ValueError
461 If an unsupported backend is specified.
462 RuntimeError
463 If nvJitLink backend is not available.
464 """
465 backend = backend.lower() 1#$%W
466 if backend != "nvjitlink": 1#$%W
467 raise ValueError(f"as_bytes() only supports 'nvjitlink' backend, got '{backend}'") 1$%
468 if not _use_nvjitlink_backend: 1#W
469 raise RuntimeError("nvJitLink backend is not available") 1#
470 return self._prepare_nvjitlink_options(as_bytes=True) 1W
473# =============================================================================
474# Private implementation: cdef inline helpers
475# =============================================================================
477cdef inline int Linker_init(Linker self, tuple object_codes, object options) except -1:
478 """Initialize a Linker instance."""
479 if len(object_codes) == 0: 1(QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
480 raise ValueError("At least one ObjectCode object must be provided") 1(
482 cdef cynvjitlink.nvJitLinkHandle c_raw_nvjitlink
483 cdef cydriver.CUlinkState c_raw_culink
484 cdef Py_ssize_t c_num_opts, i
485 cdef vector[const_char_ptr] c_str_opts
486 cdef cydriver.CUjit_option* c_drv_jit_keys_ptr
487 cdef void** c_drv_jit_values_ptr
489 self._options = options = check_or_create_options(LinkerOptions, options, "Linker options") 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
491 if _use_nvjitlink_backend: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
492 self._use_nvjitlink = True 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
493 options_bytes = options._prepare_nvjitlink_options(as_bytes=True) 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
494 c_num_opts = len(options_bytes) 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
495 c_str_opts.resize(c_num_opts) 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
496 for i in range(c_num_opts): 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
497 c_str_opts[i] = <const char*>(<bytes>options_bytes[i]) 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
498 with nogil: 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
499 HANDLE_RETURN_NVJITLINK(NULL, cynvjitlink.nvJitLinkCreate( 1QuTOqzArBsCmknDEFGHIvowxpyJVKSLltMNRcdefghiabj
500 &c_raw_nvjitlink, <uint32_t>c_num_opts, c_str_opts.data()))
501 self._nvjitlink_handle = create_nvjitlink_handle(c_raw_nvjitlink) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
502 else:
503 self._use_nvjitlink = False
504 formatted_options, option_keys = options._prepare_driver_options()
505 # Keep the formatted_options list alive: it contains bytearrays that
506 # the driver writes into via raw pointers during linking operations.
507 self._drv_log_bufs = formatted_options
508 c_num_opts = len(option_keys)
509 self._drv_jit_keys.resize(c_num_opts)
510 self._drv_jit_values.resize(c_num_opts)
511 for i in range(c_num_opts):
512 self._drv_jit_keys[i] = <cydriver.CUjit_option><int>option_keys[i]
513 val = formatted_options[i]
514 if isinstance(val, bytearray):
515 self._drv_jit_values[i] = <void*>PyByteArray_AS_STRING(val)
516 else:
517 self._drv_jit_values[i] = <void*><intptr_t>int(val)
518 c_drv_jit_keys_ptr = self._drv_jit_keys.data()
519 c_drv_jit_values_ptr = self._drv_jit_values.data()
520 try:
521 with nogil:
522 HANDLE_RETURN(cydriver.cuLinkCreate(
523 <unsigned int>c_num_opts,
524 c_drv_jit_keys_ptr,
525 c_drv_jit_values_ptr,
526 &c_raw_culink))
527 except CUDAError as e:
528 Linker_annotate_error_log(self, e)
529 raise
530 self._culink_handle = create_culink_handle(c_raw_culink)
532 for code in object_codes: 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
533 assert_type(code, ObjectCode) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
534 Linker_add_code_object(self, code) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
535 return 0 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
538cdef inline void Linker_add_code_object(Linker self, object object_code) except *:
539 """Add a single ObjectCode to the linker."""
540 data = object_code.code 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
541 cdef cynvjitlink.nvJitLinkHandle c_nvjitlink_h
542 cdef cydriver.CUlinkState c_culink_state
543 cdef cynvjitlink.nvJitLinkInputType c_nv_input_type
544 cdef cydriver.CUjitInputType c_drv_input_type
545 cdef const char* c_data_ptr
546 cdef size_t c_data_size
547 cdef const char* c_name_ptr
548 cdef const char* c_file_ptr
550 name_bytes = f"{object_code.name}".encode() 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
551 c_name_ptr = <const char*>name_bytes 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
553 input_types = _nvjitlink_input_types if self._use_nvjitlink else _driver_input_types 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
554 py_input_type = input_types.get(object_code.code_type) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
555 if py_input_type is None: 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
556 raise ValueError(f"Unknown code_type associated with ObjectCode: {object_code.code_type}")
558 if self._use_nvjitlink: 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
559 c_nvjitlink_h = as_cu(self._nvjitlink_handle) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
560 c_nv_input_type = <cynvjitlink.nvJitLinkInputType><int>py_input_type 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
561 if isinstance(data, bytes): 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
562 c_data_ptr = <const char*>(<bytes>data) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
563 c_data_size = len(data) 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
564 with nogil: 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
565 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, cynvjitlink.nvJitLinkAddData( 1QuTOqzArBsCmknDEFGHIvowxpyJKSLltMNRcdefghiabj
566 c_nvjitlink_h, c_nv_input_type, <const void*>c_data_ptr, c_data_size, c_name_ptr))
567 elif isinstance(data, str):
568 file_bytes = data.encode()
569 c_file_ptr = <const char*>file_bytes
570 with nogil:
571 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, cynvjitlink.nvJitLinkAddFile(
572 c_nvjitlink_h, c_nv_input_type, c_file_ptr))
573 else:
574 raise TypeError(f"Expected bytes or str, but got {type(data).__name__}")
575 else:
576 c_culink_state = as_cu(self._culink_handle)
577 c_drv_input_type = <cydriver.CUjitInputType><int>py_input_type
578 try:
579 if isinstance(data, bytes):
580 c_data_ptr = <const char*>(<bytes>data)
581 c_data_size = len(data)
582 with nogil:
583 HANDLE_RETURN(cydriver.cuLinkAddData(
584 c_culink_state, c_drv_input_type, <void*>c_data_ptr, c_data_size, c_name_ptr,
585 0, NULL, NULL))
586 elif isinstance(data, str):
587 file_bytes = data.encode()
588 c_file_ptr = <const char*>file_bytes
589 with nogil:
590 HANDLE_RETURN(cydriver.cuLinkAddFile(
591 c_culink_state, c_drv_input_type, c_file_ptr, 0, NULL, NULL))
592 else:
593 raise TypeError(f"Expected bytes or str, but got {type(data).__name__}")
594 except CUDAError as e:
595 Linker_annotate_error_log(self, e)
596 raise
599cdef inline object Linker_link(Linker self, str target_type):
600 """Complete linking and return the result as ObjectCode."""
601 if target_type not in ("cubin", "ptx"): 1QuOqzArBsCmknDEFGHIvowxpyJKSLltMNcdefghiabj
602 raise ValueError(f"Unsupported target type: {target_type}") 1S
604 cdef cynvjitlink.nvJitLinkHandle c_nvjitlink_h
605 cdef cydriver.CUlinkState c_culink_state
606 cdef size_t c_output_size = 0 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
607 cdef char* c_code_ptr
608 cdef void* c_cubin_out = NULL 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
610 if self._use_nvjitlink: 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
611 c_nvjitlink_h = as_cu(self._nvjitlink_handle) 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
612 with nogil: 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
613 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, cynvjitlink.nvJitLinkComplete(c_nvjitlink_h)) 1QuOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
614 if target_type == "cubin": 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
615 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
616 cynvjitlink.nvJitLinkGetLinkedCubinSize(c_nvjitlink_h, &c_output_size)) 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
617 code = bytearray(c_output_size) 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
618 c_code_ptr = <char*>(<bytearray>code) 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
619 with nogil: 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
620 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
621 cynvjitlink.nvJitLinkGetLinkedCubin(c_nvjitlink_h, c_code_ptr)) 1uOqzArBsCmknDEFGHIvowxpyJKLtMNcdefghiabj
622 else:
623 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, 1l
624 cynvjitlink.nvJitLinkGetLinkedPtxSize(c_nvjitlink_h, &c_output_size)) 1l
625 code = bytearray(c_output_size) 1l
626 c_code_ptr = <char*>(<bytearray>code) 1l
627 with nogil: 1l
628 HANDLE_RETURN_NVJITLINK(c_nvjitlink_h, 1l
629 cynvjitlink.nvJitLinkGetLinkedPtx(c_nvjitlink_h, c_code_ptr)) 1l
630 else:
631 c_culink_state = as_cu(self._culink_handle)
632 try:
633 with nogil:
634 HANDLE_RETURN(cydriver.cuLinkComplete(c_culink_state, &c_cubin_out, &c_output_size))
635 except CUDAError as e:
636 Linker_annotate_error_log(self, e)
637 raise
638 code = (<char*>c_cubin_out)[:c_output_size]
640 # Linking is complete; cache the decoded logs. cuLinkDestroy may still
641 # dereference the raw log-buffer pointers, so retain them until close().
642 self._info_log = self.get_info_log() 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
643 self._error_log = self.get_error_log() 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
645 return ObjectCode._init(bytes(code), target_type, name=self._options.name) 1uOqzArBsCmknDEFGHIvowxpyJKLltMNcdefghiabj
648cdef inline void Linker_annotate_error_log(Linker self, object e):
649 """Annotate a CUDAError with the driver linker error log."""
650 error_log = self.get_error_log()
651 if error_log:
652 e.args = (e.args[0] + f"\nLinker error log: {error_log}", *e.args[1:])
655# =============================================================================
656# Private implementation: module-level state and initialization
657# =============================================================================
659# TODO: revisit this treatment for py313t builds
660_driver = None # populated if nvJitLink cannot be used
661_inited = False
662_use_nvjitlink_backend = None # set by _decide_nvjitlink_or_driver()
664# Input type mappings populated by _lazy_init() with C-level enum ints.
665_nvjitlink_input_types = None
666_driver_input_types = None
669def _nvjitlink_has_version_symbol(nvjitlink) -> bool:
670 # This condition is equivalent to testing for version >= 12.3
671 return bool(nvjitlink._inspect_function_pointer("__nvJitLinkVersion"))
674# Note: this function is reused in the tests
675def _decide_nvjitlink_or_driver() -> bool:
676 """Return True if falling back to the cuLink* driver APIs."""
677 global _driver, _use_nvjitlink_backend
678 if _use_nvjitlink_backend is not None: 2P ) * O q z A r B s C m k n D E F G H I v o w x p y J ! ' R c d e f g h i a b j + , - . / : ; = ? @ [ ] ^ _ ` { | } ~ abbbcbdbebfb
679 return not _use_nvjitlink_backend 2P ) * O q z A r B s C m k n D E F G H I v o w x p y J R c d e f g h i a b j + , - . / : ; = ? @ [ ] ^ _ ` { | } ~ abbbcbdbebfb
681 warn_txt_common = (
682 "the driver APIs will be used instead, which do not support" 1P!'
683 " minor version compatibility or linking LTO IRs."
684 " For best results, consider upgrading to a recent version of"
685 )
687 nvjitlink_module = _optional_cuda_import( 1P!'
688 "cuda.bindings.nvjitlink",
689 probe_function=lambda module: module.version(), # probe triggers nvJitLink runtime load 1P!'
690 )
691 if nvjitlink_module is None: 1P!
692 warn_txt = f"cuda.bindings.nvjitlink is not available, therefore {warn_txt_common} cuda-bindings." 1!
693 else:
694 from cuda.bindings._internal import nvjitlink
696 if _nvjitlink_has_version_symbol(nvjitlink):
697 _use_nvjitlink_backend = True
698 return False # Use nvjitlink
699 warn_txt = (
700 f"{'nvJitLink*.dll' if sys.platform == 'win32' else 'libnvJitLink.so*'} is too old (<12.3)."
701 f" Therefore cuda.bindings.nvjitlink is not usable and {warn_txt_common} nvJitLink."
702 )
704 warn(warn_txt, stacklevel=2, category=RuntimeWarning) 1!
705 _use_nvjitlink_backend = False 1!
706 _driver = driver 1!
707 return True 1!
710def _lazy_init() -> None:
711 global _inited, _nvjitlink_input_types, _driver_input_types
712 if _inited: 1P#QuTVKSLlt$%W7UXYZ012398654MNRcdefghiabj
713 return 1P#QuTVKSLlt$%W7UXYZ012398654MNRcdefghiabj
715 _decide_nvjitlink_or_driver()
716 if _use_nvjitlink_backend:
717 _nvjitlink_input_types = {
718 "ptx": <int>cynvjitlink.NVJITLINK_INPUT_PTX,
719 "cubin": <int>cynvjitlink.NVJITLINK_INPUT_CUBIN,
720 "fatbin": <int>cynvjitlink.NVJITLINK_INPUT_FATBIN,
721 "ltoir": <int>cynvjitlink.NVJITLINK_INPUT_LTOIR,
722 "object": <int>cynvjitlink.NVJITLINK_INPUT_OBJECT,
723 "library": <int>cynvjitlink.NVJITLINK_INPUT_LIBRARY,
724 }
725 else:
726 _driver_input_types = {
727 "ptx": <int>cydriver.CU_JIT_INPUT_PTX,
728 "cubin": <int>cydriver.CU_JIT_INPUT_CUBIN,
729 "fatbin": <int>cydriver.CU_JIT_INPUT_FATBINARY,
730 "object": <int>cydriver.CU_JIT_INPUT_OBJECT,
731 "library": <int>cydriver.CU_JIT_INPUT_LIBRARY,
732 }
733 _inited = True