44
55from __future__ import annotations
66
7+ from libc.stdint cimport uintptr_t
8+
9+ # TODO: how about cuda.bindings < 12.6.2?
10+ from cuda.bindings cimport cydriver
11+
712from cuda.core.experimental._utils.cuda_utils cimport (
813 _check_driver_error as raise_if_driver_error,
914 check_or_create_options,
1015)
16+
1117import sys
1218
1319import cython
@@ -59,7 +65,7 @@ class IsStreamT(Protocol):
5965 ...
6066
6167
62- def _try_to_get_stream_ptr(obj: IsStreamT ):
68+ cdef cydriver.CUstream _try_to_get_stream_ptr(obj: IsStreamT ) except* :
6369 try:
6470 cuda_stream_attr = obj.__cuda_stream__
6571 except AttributeError:
@@ -86,7 +92,7 @@ def _try_to_get_stream_ptr(obj: IsStreamT):
8692 raise RuntimeError (
8793 f" The first element of the sequence returned by obj.__cuda_stream__ must be 0, got {repr(info[0])}"
8894 )
89- return driver .CUstream(info[1 ])
95+ return < cydriver .CUstream>< uintptr_t > (info[1 ])
9096
9197
9298cdef class Stream:
@@ -108,14 +114,17 @@ cdef class Stream:
108114 """
109115
110116 cdef:
111- object _handle
117+ cydriver.CUstream _handle
112118 object _owner
113119 object _builtin
114120 object _nonblocking
115121 object _priority
116122 object _device_id
117123 object _ctx_handle
118124
125+ def __cinit__ (self , *args , **kwargs ):
126+ self ._handle = < cydriver.CUstream> (NULL )
127+
119128 def __init__ (self , *args , **kwargs ):
120129 raise RuntimeError (
121130 " Stream objects cannot be instantiated directly. "
@@ -125,7 +134,7 @@ cdef class Stream:
125134 @classmethod
126135 def _legacy_default (cls ):
127136 cdef Stream self = Stream.__new__ (cls )
128- self ._handle = driver .CUstream(driver .CU_STREAM_LEGACY)
137+ self ._handle = < cydriver .CUstream> (cydriver .CU_STREAM_LEGACY)
129138 self ._owner = None
130139 self ._builtin = True
131140 self ._nonblocking = None # delayed
@@ -137,7 +146,7 @@ cdef class Stream:
137146 @classmethod
138147 def _per_thread_default (cls ):
139148 cdef Stream self = Stream.__new__ (cls )
140- self ._handle = driver .CUstream(driver .CU_STREAM_PER_THREAD)
149+ self ._handle = < cydriver .CUstream> (cydriver .CU_STREAM_PER_THREAD)
141150 self ._owner = None
142151 self ._builtin = True
143152 self ._nonblocking = None # delayed
@@ -149,7 +158,6 @@ cdef class Stream:
149158 @classmethod
150159 def _init (cls , obj: Optional[IsStreamT] = None , options = None , device_id: int = None ):
151160 cdef Stream self = Stream.__new__ (cls )
152- self ._handle = None
153161 self ._owner = None
154162 self ._builtin = False
155163
@@ -169,16 +177,20 @@ cdef class Stream:
169177 nonblocking = opts.nonblocking
170178 priority = opts.priority
171179
172- flags = driver.CUstream_flags.CU_STREAM_NON_BLOCKING if nonblocking else driver.CUstream_flags.CU_STREAM_DEFAULT
173- err, high, low = driver.cuCtxGetStreamPriorityRange()
174- raise_if_driver_error(err)
180+ flags = cydriver.CUstream_flags.CU_STREAM_NON_BLOCKING if nonblocking else cydriver.CUstream_flags.CU_STREAM_DEFAULT
181+ # TODO: use HANDLE_RETURN
182+ cdef int high, low
183+ err = cydriver.cuCtxGetStreamPriorityRange(& high, & low)
175184 if priority is not None :
176185 if not (low <= priority <= high):
177186 raise ValueError (f" {priority=} is out of range {[low, high]}" )
178187 else :
179188 priority = high
180189
181- self ._handle = handle_return(driver.cuStreamCreateWithPriority(flags, priority))
190+ cdef cydriver.CUstream s
191+ # TODO: add HANDLE_RETURN macro to check driver error code?
192+ err = cydriver.cuStreamCreateWithPriority(& s, flags, priority)
193+ self ._handle = s
182194 self ._owner = None
183195 self ._nonblocking = nonblocking
184196 self ._priority = priority
@@ -195,10 +207,11 @@ cdef class Stream:
195207
196208 if self ._owner is None :
197209 if self ._handle and not self ._builtin:
198- handle_return(driver.cuStreamDestroy(self ._handle))
210+ # TODO: use HANDLE_RETURN
211+ err = cydriver.cuStreamDestroy(self ._handle)
199212 else :
200213 self ._owner = None
201- self ._handle = None
214+ self ._handle = < cydriver.CUstream > ( NULL )
202215
203216 cpdef close(self ):
204217 """ Destroy the stream.
@@ -222,14 +235,16 @@ cdef class Stream:
222235 This handle is a Python object. To get the memory address of the underlying C
223236 handle , call ``int(Stream.handle )``.
224237 """
225- return self._handle
238+ return driver.CUstream(<uintptr_t><void*>( self._handle ))
226239
227240 @property
228241 def is_nonblocking(self ) -> bool:
229242 """Return True if this is a nonblocking stream , otherwise False."""
243+ cdef unsigned int flags
230244 if self._nonblocking is None:
231- flag = handle_return(driver.cuStreamGetFlags(self ._handle))
232- if flag == driver.CUstream_flags.CU_STREAM_NON_BLOCKING:
245+ # TODO: switch to HANDLE_RETURN
246+ err = cydriver.cuStreamGetFlags(self ._handle, & flags)
247+ if flags & cydriver.CUstream_flags.CU_STREAM_NON_BLOCKING:
233248 self._nonblocking = True
234249 else:
235250 self._nonblocking = False
@@ -238,14 +253,17 @@ cdef class Stream:
238253 @property
239254 def priority(self ) -> int:
240255 """Return the stream priority."""
256+ cdef int prio
241257 if self._priority is None:
242- prio = handle_return(driver.cuStreamGetPriority(self ._handle))
258+ # TODO: switch to HANDLE_RETURN
259+ err = cydriver.cuStreamGetPriority(self ._handle, & prio)
243260 self._priority = prio
244261 return self._priority
245262
246263 def sync(self ):
247264 """ Synchronize the stream."""
248- handle_return(driver.cuStreamSynchronize(self ._handle))
265+ # TODO: switch to HANDLE_RETURN
266+ err = cydriver.cuStreamSynchronize(self ._handle)
249267
250268 def record (self , event: Event = None , options: EventOptions = None ) -> Event:
251269 """Record an event onto the stream.
@@ -272,8 +290,9 @@ cdef class Stream:
272290 if event is None:
273291 self._get_device_and_context()
274292 event = Event._init(self ._device_id, self ._ctx_handle, options)
275- err , = driver.cuEventRecord(event.handle , self._handle )
276- raise_if_driver_error(err )
293+ # TODO: switch to HANDLE_RETURN
294+ # TODO: revisit after Event is cythonized
295+ err = cydriver.cuEventRecord(< cydriver.CUevent>< uintptr_t> (event.handle), self ._handle)
277296 return event
278297
279298 def wait(self , event_or_stream: Union[Event , Stream]):
@@ -286,28 +305,35 @@ cdef class Stream:
286305 on the stream and then waiting on it.
287306
288307 """
308+ cdef cydriver.CUevent event
309+ cdef cydriver.CUstream stream
310+ cdef bint discard_event
311+
289312 if isinstance (event_or_stream, Event):
290- event = event_or_stream.handle
313+ event = < cydriver.CUevent >< uintptr_t > ( event_or_stream.handle)
291314 discard_event = False
292315 else :
293316 if isinstance (event_or_stream, Stream):
294- stream = event_or_stream
317+ stream = < cydriver.CUstream >< uintptr_t > ( event_or_stream.handle)
295318 else :
296319 try :
297- stream = Stream._init(obj = event_or_stream)
320+ s = Stream._init(obj = event_or_stream)
298321 except Exception as e:
299322 raise ValueError (
300323 " only an Event, Stream, or object supporting __cuda_stream__ can be waited,"
301324 f" got {type(event_or_stream)}"
302325 ) from e
303- event = handle_return(driver.cuEventCreate(driver.CUevent_flags.CU_EVENT_DISABLE_TIMING))
304- handle_return(driver.cuEventRecord(event, stream.handle))
326+ stream = < cydriver.CUstream>< uintptr_t> (s.handle)
327+ # TODO: switch to HANDLE_RETURN
328+ err = cydriver.cuEventCreate(& event, cydriver.CUevent_flags.CU_EVENT_DISABLE_TIMING)
329+ err = cydriver.cuEventRecord(event, stream)
305330 discard_event = True
306331
307332 # TODO: support flags other than 0?
308- handle_return(driver.cuStreamWaitEvent(self ._handle, event, 0 ))
333+ # TODO: switch to HANDLE_RETURN
334+ err = cydriver.cuStreamWaitEvent(self ._handle, event, 0 )
309335 if discard_event:
310- handle_return(driver .cuEventDestroy(event) )
336+ err = cydriver .cuEventDestroy(event)
311337
312338 @property
313339 def device (self ) -> Device:
@@ -325,9 +351,12 @@ cdef class Stream:
325351 return Device(self._device_id )
326352
327353 cdef int _get_context(Stream self ) except?-1:
354+ # TODO: consider making self._ctx_handle typed?
355+ cdef cydriver.CUcontext ctx
328356 if self._ctx_handle is None:
329- err , self._ctx_handle = driver.cuStreamGetCtx(self ._handle)
330- raise_if_driver_error(err )
357+ # TODO: switch to HANDLE_RETURN
358+ err = cydriver.cuStreamGetCtx(self ._handle, & ctx)
359+ self._ctx_handle = driver.CUcontext(< uintptr_t> ctx)
331360 return 0
332361
333362 cdef int _get_device_and_context(Stream self ) except?-1:
0 commit comments