Skip to content

Commit 059829c

Browse files
committed
[DONT MERGE] Make faster
1 parent 89909f3 commit 059829c

5 files changed

Lines changed: 159 additions & 97 deletions

File tree

cuda_core/cuda/core/experimental/_stream.py renamed to cuda_core/cuda/core/experimental/_stream.pyx

Lines changed: 75 additions & 58 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,11 @@
44

55
from __future__ import annotations
66

7+
from cython import NULL
8+
from libcpp cimport bool
9+
from libc.stdint cimport intptr_t
10+
11+
712
import os
813
import warnings
914
import weakref
@@ -13,6 +18,7 @@
1318
if TYPE_CHECKING:
1419
import cuda.bindings
1520
from cuda.core.experimental._device import Device
21+
1622
from cuda.core.experimental._context import Context
1723
from cuda.core.experimental._event import Event, EventOptions
1824
from cuda.core.experimental._utils.clear_error_support import assert_type
@@ -24,6 +30,10 @@
2430
runtime,
2531
)
2632

33+
from cuda.bindings cimport cydriver as cdriver
34+
from cuda.bindings cimport cyruntime as cruntime
35+
from cuda.core.experimental._utils cimport _error_utils
36+
2737

2838
@dataclass
2939
class StreamOptions:
@@ -43,7 +53,7 @@ class StreamOptions:
4353
priority: Optional[int] = None
4454

4555

46-
class Stream:
56+
cdef class Stream:
4757
"""Represent a queue of GPU operations that are executed in a specific order.
4858
4959
Applications use streams to control the order of execution for
@@ -61,55 +71,42 @@ class Stream:
6171
6272
"""
6373

64-
class _MembersNeededForFinalize:
65-
__slots__ = ("handle", "owner", "builtin")
66-
67-
def __init__(self, stream_obj, handle, owner, builtin):
68-
self.handle = handle
69-
self.owner = owner
70-
self.builtin = builtin
71-
weakref.finalize(stream_obj, self.close)
72-
73-
def close(self):
74-
if self.owner is None:
75-
if self.handle and not self.builtin:
76-
handle_return(driver.cuStreamDestroy(self.handle))
77-
else:
78-
self.owner = None
79-
self.handle = None
80-
81-
def __new__(self, *args, **kwargs):
74+
def __init__(self, *args, **kwargs):
8275
raise RuntimeError(
8376
"Stream objects cannot be instantiated directly. "
8477
"Please use Device APIs (create_stream) or other Stream APIs (from_handle)."
8578
)
8679

87-
__slots__ = ("__weakref__", "_mnff", "_nonblocking", "_priority", "_device_id", "_ctx_handle")
80+
__slots__ = ("_handle", "_nonblocking", "_priority", "_device_id", "_ctx_handle")
81+
82+
cdef cdriver.CUstream _handle
83+
cdef bool _nonblocking
84+
cdef int _priority
85+
cdef int _device_id
86+
cdef object _ctx_handle # Python object
87+
8888

8989
@classmethod
9090
def _legacy_default(cls):
91-
self = super().__new__(cls)
92-
self._mnff = Stream._MembersNeededForFinalize(self, driver.CUstream(driver.CU_STREAM_LEGACY), None, True)
93-
self._nonblocking = None # delayed
94-
self._priority = None # delayed
95-
self._device_id = None # delayed
91+
cdef Stream self = cls.__new__(cls) # safe allocator
92+
self._nonblocking = True # delayed
93+
self._priority = 0 # delayed
94+
self._device_id = -1 # delayed
9695
self._ctx_handle = None # delayed
9796
return self
9897

9998
@classmethod
10099
def _per_thread_default(cls):
101-
self = super().__new__(cls)
102-
self._mnff = Stream._MembersNeededForFinalize(self, driver.CUstream(driver.CU_STREAM_PER_THREAD), None, True)
103-
self._nonblocking = None # delayed
104-
self._priority = None # delayed
105-
self._device_id = None # delayed
100+
cdef Stream self = cls.__new__(cls) # safe allocator
101+
self._nonblocking = True # delayed
102+
self._priority = 0 # delayed
103+
self._device_id = -1 # delayed
106104
self._ctx_handle = None # delayed
107105
return self
108106

109107
@classmethod
110108
def _init(cls, obj=None, *, options: Optional[StreamOptions] = None):
111-
self = super().__new__(cls)
112-
self._mnff = Stream._MembersNeededForFinalize(self, None, None, False)
109+
cdef Stream self = cls.__new__(cls)
113110

114111
if obj is not None and options is not None:
115112
raise ValueError("obj and options cannot be both specified")
@@ -142,46 +139,66 @@ def _init(cls, obj=None, *, options: Optional[StreamOptions] = None):
142139
f"The first element of the sequence returned by obj.__cuda_stream__ must be 0, got {repr(info[0])}"
143140
)
144141

145-
self._mnff.handle = driver.CUstream(info[1])
146142
# TODO: check if obj is created under the current context/device
147-
self._mnff.owner = obj
148-
self._nonblocking = None # delayed
149-
self._priority = None # delayed
150-
self._device_id = None # delayed
143+
self._nonblocking = True # delayed
144+
self._priority = 0 # delayed
145+
self._device_id = -1 # delayed
151146
self._ctx_handle = None # delayed
152147
return self
153148

154-
options = check_or_create_options(StreamOptions, options, "Stream options")
155-
nonblocking = options.nonblocking
156-
priority = options.priority
149+
# options = check_or_create_options(StreamOptions, options, "Stream options")
150+
cdef int high, low = 0
151+
cdef cruntime.cudaError_t r_err = cruntime.cudaDeviceGetStreamPriorityRange(&high, &low)
152+
_error_utils._check_runtime_error(r_err)
153+
#high, low = result[1:]
154+
#high, low = handle_return(runtime.cudaDeviceGetStreamPriorityRange())
155+
156+
cdef bool nonblocking = False
157+
cdef int priority = high
158+
if options is not None:
159+
nonblocking = options.nonblocking
160+
priority = options.priority if options.priority is not None else priority
161+
162+
cdef flags = cdriver.CUstream_flags.CU_STREAM_NON_BLOCKING if nonblocking else cdriver.CUstream_flags.CU_STREAM_DEFAULT
157163

158-
flags = driver.CUstream_flags.CU_STREAM_NON_BLOCKING if nonblocking else driver.CUstream_flags.CU_STREAM_DEFAULT
159164

160-
high, low = handle_return(runtime.cudaDeviceGetStreamPriorityRange())
161165
if priority is not None:
162166
if not (low <= priority <= high):
163167
raise ValueError(f"{priority=} is out of range {[low, high]}")
164-
else:
165-
priority = high
166168

167-
self._mnff.handle = handle_return(driver.cuStreamCreateWithPriority(flags, priority))
168-
self._mnff.owner = None
169+
#cdef cdriver.CUstream handler = cdriver.CUstream()
170+
cdef cdriver.CUresult c_err = cdriver.cuStreamCreateWithPriority(&self._handle, flags, priority)
171+
_error_utils._check_driver_error(c_err)
172+
173+
#self._handle = handler
174+
169175
self._nonblocking = nonblocking
170176
self._priority = priority
171177
# don't defer this because we will have to pay a cost for context
172178
# switch later
173-
self._device_id = int(handle_return(driver.cuCtxGetDevice()))
179+
cdef int device_id = -1
180+
err = cdriver.cuCtxGetDevice(&device_id)
181+
_error_utils._check_driver_error(err)
182+
self._device_id = device_id
183+
#self._device_id = int(handle_return(driver.cuCtxGetDevice()))
174184
self._ctx_handle = None # delayed
175185
return self
176186

187+
def __dealloc__(self):
188+
if self._handle:
189+
c_err = cdriver.cuStreamDestroy(self._handle)
190+
_error_utils._check_driver_error(c_err)
191+
177192
def close(self):
178193
"""Destroy the stream.
179194
180195
Destroy the stream if we own it. Borrowed foreign stream
181196
object will instead have their references released.
182197
183198
"""
184-
self._mnff.close()
199+
if self.handle:
200+
handle_return(driver.cuStreamDestroy(self.handle))
201+
self._handle = NULL
185202

186203
def __cuda_stream__(self) -> Tuple[int, int]:
187204
"""Return an instance of a __cuda_stream__ protocol."""
@@ -193,16 +210,16 @@ def handle(self) -> cuda.bindings.driver.CUstream:
193210

194211
.. caution::
195212

196-
This handle is a Python object. To get the memory address of the underlying C
197-
handle, call ``int(Stream.handle)``.
213+
This handle is a Python object. Representing the memory
214+
address of the underlying C object.
198215
"""
199-
return self._mnff.handle
216+
return <intptr_t>self._handle
200217

201218
@property
202219
def is_nonblocking(self) -> bool:
203220
"""Return True if this is a nonblocking stream, otherwise False."""
204221
if self._nonblocking is None:
205-
flag = handle_return(driver.cuStreamGetFlags(self._mnff.handle))
222+
flag = handle_return(driver.cuStreamGetFlags(self.handle))
206223
if flag == driver.CUstream_flags.CU_STREAM_NON_BLOCKING:
207224
self._nonblocking = True
208225
else:
@@ -213,13 +230,13 @@ def is_nonblocking(self) -> bool:
213230
def priority(self) -> int:
214231
"""Return the stream priority."""
215232
if self._priority is None:
216-
prio = handle_return(driver.cuStreamGetPriority(self._mnff.handle))
233+
prio = handle_return(driver.cuStreamGetPriority(self.handle))
217234
self._priority = prio
218235
return self._priority
219236

220237
def sync(self):
221238
"""Synchronize the stream."""
222-
handle_return(driver.cuStreamSynchronize(self._mnff.handle))
239+
handle_return(driver.cuStreamSynchronize(self.handle))
223240

224241
def record(self, event: Event = None, options: EventOptions = None) -> Event:
225242
"""Record an event onto the stream.
@@ -246,7 +263,7 @@ def record(self, event: Event = None, options: EventOptions = None) -> Event:
246263
if event is None:
247264
event = Event._init(self._device_id, self._ctx_handle, options)
248265
assert_type(event, Event)
249-
handle_return(driver.cuEventRecord(event.handle, self._mnff.handle))
266+
handle_return(driver.cuEventRecord(event.handle, self.handle))
250267
return event
251268

252269
def wait(self, event_or_stream: Union[Event, Stream]):
@@ -278,7 +295,7 @@ def wait(self, event_or_stream: Union[Event, Stream]):
278295
discard_event = True
279296

280297
# TODO: support flags other than 0?
281-
handle_return(driver.cuStreamWaitEvent(self._mnff.handle, event, 0))
298+
handle_return(driver.cuStreamWaitEvent(self.handle, event, 0))
282299
if discard_event:
283300
handle_return(driver.cuEventDestroy(event))
284301

@@ -298,15 +315,15 @@ def device(self) -> Device:
298315
if self._device_id is None:
299316
# Get the stream context first
300317
if self._ctx_handle is None:
301-
self._ctx_handle = handle_return(driver.cuStreamGetCtx(self._mnff.handle))
318+
self._ctx_handle = handle_return(driver.cuStreamGetCtx(self.handle))
302319
self._device_id = get_device_from_ctx(self._ctx_handle)
303320
return Device(self._device_id)
304321

305322
@property
306323
def context(self) -> Context:
307324
"""Return the :obj:`~_context.Context` associated with this stream."""
308325
if self._ctx_handle is None:
309-
self._ctx_handle = handle_return(driver.cuStreamGetCtx(self._mnff.handle))
326+
self._ctx_handle = handle_return(driver.cuStreamGetCtx(self.handle))
310327
if self._device_id is None:
311328
self._device_id = get_device_from_ctx(self._ctx_handle)
312329
return Context._from_ctx(self._ctx_handle, self._device_id)
Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
from cuda.bindings cimport cydriver as driver
2+
from cuda.bindings cimport cyruntime as runtime
3+
4+
cpdef _check_driver_error(driver.CUresult result)
5+
cpdef _check_runtime_error(runtime.cudaError_t error)
Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,56 @@
1+
from cython cimport NULL
2+
from libc.string cimport strlen
3+
4+
import cuda.bindings
5+
from cuda.bindings import driver, runtime
6+
7+
from cuda.bindings cimport cydriver as cdriver
8+
from cuda.bindings cimport cyruntime as cruntime
9+
10+
from cuda.core.experimental._utils.driver_cu_result_explanations import DRIVER_CU_RESULT_EXPLANATIONS
11+
from cuda.core.experimental._utils.runtime_cuda_error_explanations import RUNTIME_CUDA_ERROR_EXPLANATIONS
12+
13+
14+
class CUDAError(Exception):
15+
pass
16+
17+
18+
19+
cpdef _check_driver_error(cdriver.CUresult error):
20+
if error == cdriver.CUresult.CUDA_SUCCESS:
21+
return
22+
cdef const char *c_name = NULL
23+
name_err = driver.cuGetErrorName(error, &c_name)
24+
if name_err != driver.CUresult.CUDA_SUCCESS:
25+
raise CUDAError(f"UNEXPECTED ERROR CODE: {error}")
26+
name: str = (<char *>c_name)[:strlen(c_name)].decode()
27+
#name = name.decode()
28+
expl = DRIVER_CU_RESULT_EXPLANATIONS.get(int(error))
29+
if expl is not None:
30+
raise CUDAError(f"{name}: {expl}")
31+
cdef const char *c_desc = NULL
32+
desc_err = driver.cuGetErrorString(error, &c_desc)
33+
if desc_err != driver.CUresult.CUDA_SUCCESS:
34+
raise CUDAError(f"{name}")
35+
desc: str = (<char *>c_desc)[:strlen(c_desc)].decode()
36+
#desc = desc.decode()
37+
raise CUDAError(f"{name}: {desc}")
38+
39+
40+
cpdef _check_runtime_error(cruntime.cudaError_t error):
41+
if error == cruntime.cudaError_t.cudaSuccess:
42+
return
43+
name_err, name = runtime.cudaGetErrorName(error)
44+
if name_err != runtime.cudaError_t.cudaSuccess:
45+
raise CUDAError(f"UNEXPECTED ERROR CODE: {error}")
46+
name = name.decode()
47+
expl = RUNTIME_CUDA_ERROR_EXPLANATIONS.get(int(error))
48+
if expl is not None:
49+
raise CUDAError(f"{name}: {expl}")
50+
desc_err, desc = runtime.cudaGetErrorString(error)
51+
if desc_err != runtime.cudaError_t.cudaSuccess:
52+
raise CUDAError(f"{name}")
53+
desc = desc.decode()
54+
raise CUDAError(f"{name}: {desc}")
55+
56+

cuda_core/cuda/core/experimental/_utils/cuda_utils.py

Lines changed: 1 addition & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -15,14 +15,11 @@
1515
from cuda import cudart as runtime
1616
from cuda import nvrtc
1717

18+
from cuda.core.experimental._utils._error_utils import CUDAError, _check_driver_error, _check_runtime_error
1819
from cuda.core.experimental._utils.driver_cu_result_explanations import DRIVER_CU_RESULT_EXPLANATIONS
1920
from cuda.core.experimental._utils.runtime_cuda_error_explanations import RUNTIME_CUDA_ERROR_EXPLANATIONS
2021

2122

22-
class CUDAError(Exception):
23-
pass
24-
25-
2623
class NVRTCError(CUDAError):
2724
pass
2825

@@ -48,40 +45,6 @@ def cast_to_3_tuple(label, cfg):
4845
return cfg + (1,) * (3 - len(cfg))
4946

5047

51-
def _check_driver_error(error):
52-
if error == driver.CUresult.CUDA_SUCCESS:
53-
return
54-
name_err, name = driver.cuGetErrorName(error)
55-
if name_err != driver.CUresult.CUDA_SUCCESS:
56-
raise CUDAError(f"UNEXPECTED ERROR CODE: {error}")
57-
name = name.decode()
58-
expl = DRIVER_CU_RESULT_EXPLANATIONS.get(int(error))
59-
if expl is not None:
60-
raise CUDAError(f"{name}: {expl}")
61-
desc_err, desc = driver.cuGetErrorString(error)
62-
if desc_err != driver.CUresult.CUDA_SUCCESS:
63-
raise CUDAError(f"{name}")
64-
desc = desc.decode()
65-
raise CUDAError(f"{name}: {desc}")
66-
67-
68-
def _check_runtime_error(error):
69-
if error == runtime.cudaError_t.cudaSuccess:
70-
return
71-
name_err, name = runtime.cudaGetErrorName(error)
72-
if name_err != runtime.cudaError_t.cudaSuccess:
73-
raise CUDAError(f"UNEXPECTED ERROR CODE: {error}")
74-
name = name.decode()
75-
expl = RUNTIME_CUDA_ERROR_EXPLANATIONS.get(int(error))
76-
if expl is not None:
77-
raise CUDAError(f"{name}: {expl}")
78-
desc_err, desc = runtime.cudaGetErrorString(error)
79-
if desc_err != runtime.cudaError_t.cudaSuccess:
80-
raise CUDAError(f"{name}")
81-
desc = desc.decode()
82-
raise CUDAError(f"{name}: {desc}")
83-
84-
8548
def _check_error(error, handle=None):
8649
if isinstance(error, driver.CUresult):
8750
_check_driver_error(error)

0 commit comments

Comments
 (0)