44
55from __future__ import annotations
66
7+ from cython import NULL
8+ from libcpp cimport bool
9+ from libc.stdint cimport intptr_t
10+
11+
712import os
813import warnings
914import weakref
1318if TYPE_CHECKING:
1419 import cuda.bindings
1520 from cuda.core.experimental._device import Device
21+
1622from cuda.core.experimental._context import Context
1723from cuda.core.experimental._event import Event, EventOptions
1824from cuda.core.experimental._utils.clear_error_support import assert_type
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
2939class 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 )
0 commit comments