diff --git a/kernel_tuner/utils/nvcuda.py b/kernel_tuner/utils/nvcuda.py index 83dca98a..25eddd00 100644 --- a/kernel_tuner/utils/nvcuda.py +++ b/kernel_tuner/utils/nvcuda.py @@ -48,6 +48,9 @@ def cuda_error_check(error): """Checking the status of CUDA calls using the NVIDIA cuda-python backend.""" + # cuda-python always returns a tuple, also for calls that only return a status, e.g. (CUresult,) + if isinstance(error, tuple) and len(error) == 1: + error = error[0] if isinstance(error, driver.CUresult): if error != driver.CUresult.CUDA_SUCCESS: _, name = driver.cuGetErrorName(error) @@ -60,6 +63,8 @@ def cuda_error_check(error): if error != nvrtc.nvrtcResult.NVRTC_SUCCESS: _, desc = nvrtc.nvrtcGetErrorString(error) raise RuntimeError(f"NVRTC error: {desc.decode()}") + else: + raise TypeError(f"Unable to check CUDA status of unexpected type {type(error).__name__}: {error!r}") def _check(call_result):