44
55from __future__ import annotations
66
7+ from cuda.core.experimental._utils.cuda_utils cimport (
8+ _check_driver_error as raise_if_driver_error,
9+ check_or_create_options,
10+ )
11+
712import os
813import warnings
9- import weakref
1014from dataclasses import dataclass
1115from typing import TYPE_CHECKING, Optional, Protocol, Tuple, Union
1216
1822from cuda.core.experimental._graph import GraphBuilder
1923from cuda.core.experimental._utils.clear_error_support import assert_type
2024from cuda.core.experimental._utils.cuda_utils import (
21- check_or_create_options ,
2225 driver,
2326 get_device_from_ctx,
2427 handle_return,
25- runtime ,
2628)
2729
2830
2931@dataclass
30- class StreamOptions :
32+ cdef class StreamOptions:
3133 """ Customizable :obj:`~_stream.Stream` options.
3234
3335 Attributes
@@ -85,7 +87,7 @@ def _try_to_get_stream_ptr(obj: IsStreamT):
8587 return driver.CUstream(info[1 ])
8688
8789
88- class Stream :
90+ cdef class Stream:
8991 """ Represent a queue of GPU operations that are executed in a specific order.
9092
9193 Applications use streams to control the order of execution for
@@ -103,35 +105,27 @@ class Stream:
103105
104106 """
105107
106- class _MembersNeededForFinalize :
107- __slots__ = ( "handle" , "owner" , "builtin" )
108-
109- def __init__ ( self , stream_obj , handle , owner , builtin ):
110- self . handle = handle
111- self . owner = owner
112- self . builtin = builtin
113- weakref . finalize ( stream_obj , self . close )
108+ cdef :
109+ object _handle
110+ object _owner
111+ object _builtin
112+ object _nonblocking
113+ object _priority
114+ object _device_id
115+ object _ctx_handle
114116
115- def close (self ):
116- if self .owner is None :
117- if self .handle and not self .builtin :
118- handle_return (driver .cuStreamDestroy (self .handle ))
119- else :
120- self .owner = None
121- self .handle = None
122-
123- def __new__ (self , * args , ** kwargs ):
117+ def __init__ (self , *args , **kwargs ):
124118 raise RuntimeError (
125119 " Stream objects cannot be instantiated directly. "
126120 " Please use Device APIs (create_stream) or other Stream APIs (from_handle)."
127121 )
128122
129- __slots__ = ("__weakref__" , "_mnff" , "_nonblocking" , "_priority" , "_device_id" , "_ctx_handle" )
130-
131123 @classmethod
132124 def _legacy_default (cls ):
133- self = super ().__new__ (cls )
134- self ._mnff = Stream ._MembersNeededForFinalize (self , driver .CUstream (driver .CU_STREAM_LEGACY ), None , True )
125+ cdef Stream self = Stream.__new__ (Stream)
126+ self ._handle = driver.CUstream(driver.CU_STREAM_LEGACY)
127+ self ._owner = None
128+ self ._builtin = True
135129 self ._nonblocking = None # delayed
136130 self ._priority = None # delayed
137131 self ._device_id = None # delayed
@@ -140,66 +134,76 @@ def _legacy_default(cls):
140134
141135 @classmethod
142136 def _per_thread_default (cls ):
143- self = super ().__new__ (cls )
144- self ._mnff = Stream ._MembersNeededForFinalize (self , driver .CUstream (driver .CU_STREAM_PER_THREAD ), None , True )
137+ cdef Stream self = Stream.__new__ (Stream)
138+ self ._handle = driver.CUstream(driver.CU_STREAM_PER_THREAD)
139+ self ._owner = None
140+ self ._builtin = True
145141 self ._nonblocking = None # delayed
146142 self ._priority = None # delayed
147143 self ._device_id = None # delayed
148144 self ._ctx_handle = None # delayed
149145 return self
150146
151147 @classmethod
152- def _init (cls , obj : Optional [IsStreamT ] = None , * , options : Optional [StreamOptions ] = None ):
153- self = super ().__new__ (cls )
154- self ._mnff = Stream ._MembersNeededForFinalize (self , None , None , False )
148+ def _init (cls , obj: Optional[IsStreamT] = None , *, options: Optional[StreamOptions] = None , device_id: int = None ):
149+ cdef Stream self = Stream.__new__ (Stream)
150+ self ._handle = None
151+ self ._owner = None
152+ self ._builtin = False
155153
156154 if obj is not None and options is not None :
157155 raise ValueError (" obj and options cannot be both specified" )
158156 if obj is not None :
159- self ._mnff . handle = _try_to_get_stream_ptr (obj )
157+ self ._handle = _try_to_get_stream_ptr(obj)
160158 # TODO: check if obj is created under the current context/device
161- self ._mnff . owner = obj
159+ self ._owner = obj
162160 self ._nonblocking = None # delayed
163161 self ._priority = None # delayed
164162 self ._device_id = None # delayed
165163 self ._ctx_handle = None # delayed
166164 return self
167165
168- options = check_or_create_options (StreamOptions , options , "Stream options" )
169- nonblocking = options .nonblocking
170- priority = options .priority
166+ cdef StreamOptions opts = check_or_create_options(StreamOptions, options, " Stream options" )
167+ nonblocking = opts .nonblocking
168+ priority = opts .priority
171169
172170 flags = driver.CUstream_flags.CU_STREAM_NON_BLOCKING if nonblocking else driver.CUstream_flags.CU_STREAM_DEFAULT
173-
174- high , low = handle_return ( runtime . cudaDeviceGetStreamPriorityRange () )
171+ err, high, low = driver.cuCtxGetStreamPriorityRange()
172+ raise_if_driver_error(err )
175173 if priority is not None :
176174 if not (low <= priority <= high):
177175 raise ValueError (f" {priority=} is out of range {[low, high]}" )
178176 else :
179177 priority = high
180178
181- self ._mnff . handle = handle_return (driver .cuStreamCreateWithPriority (flags , priority ))
182- self ._mnff . owner = None
179+ self ._handle = handle_return(driver.cuStreamCreateWithPriority(flags, priority))
180+ self ._owner = None
183181 self ._nonblocking = nonblocking
184182 self ._priority = priority
185- # don't defer this because we will have to pay a cost for context
186- # switch later
187- self ._device_id = int (handle_return (driver .cuCtxGetDevice ()))
183+ self ._device_id = device_id
188184 self ._ctx_handle = None # delayed
189185 return self
190186
191- def close (self ):
187+ def __del__ (self ):
188+ self .close()
189+
190+ cpdef close(self ):
192191 """ Destroy the stream.
193192
194193 Destroy the stream if we own it. Borrowed foreign stream
195194 object will instead have their references released.
196195
197196 """
198- self ._mnff .close ()
197+ if self ._owner is None :
198+ if self ._handle and not self ._builtin:
199+ handle_return(driver.cuStreamDestroy(self ._handle))
200+ else :
201+ self ._owner = None
202+ self ._handle = None
199203
200204 def __cuda_stream__ (self ) -> Tuple[int , int]:
201205 """Return an instance of a __cuda_stream__ protocol."""
202- return (0 , self .handle )
206+ return (0, int( self.handle ) )
203207
204208 @property
205209 def handle(self ) -> cuda.bindings.driver.CUstream:
@@ -210,13 +214,13 @@ def handle(self) -> cuda.bindings.driver.CUstream:
210214 This handle is a Python object. To get the memory address of the underlying C
211215 handle , call ``int(Stream.handle )``.
212216 """
213- return self ._mnff . handle
217+ return self._handle
214218
215219 @property
216220 def is_nonblocking(self ) -> bool:
217221 """Return True if this is a nonblocking stream , otherwise False."""
218222 if self._nonblocking is None:
219- flag = handle_return (driver .cuStreamGetFlags (self ._mnff . handle ))
223+ flag = handle_return(driver.cuStreamGetFlags(self ._handle ))
220224 if flag == driver.CUstream_flags.CU_STREAM_NON_BLOCKING:
221225 self._nonblocking = True
222226 else:
@@ -227,13 +231,13 @@ def is_nonblocking(self) -> bool:
227231 def priority(self ) -> int:
228232 """Return the stream priority."""
229233 if self._priority is None:
230- prio = handle_return (driver .cuStreamGetPriority (self ._mnff . handle ))
234+ prio = handle_return(driver.cuStreamGetPriority(self ._handle ))
231235 self._priority = prio
232236 return self._priority
233237
234238 def sync(self ):
235239 """ Synchronize the stream."""
236- handle_return (driver .cuStreamSynchronize (self ._mnff . handle ))
240+ handle_return(driver.cuStreamSynchronize(self ._handle ))
237241
238242 def record (self , event: Event = None , options: EventOptions = None ) -> Event:
239243 """Record an event onto the stream.
@@ -259,8 +263,8 @@ def record(self, event: Event = None, options: EventOptions = None) -> Event:
259263 # and CU_EVENT_RECORD_EXTERNAL , can be set in EventOptions.
260264 if event is None:
261265 event = Event._init(self ._device_id, self ._ctx_handle, options)
262- assert_type (event , Event )
263- handle_return ( driver . cuEventRecord ( event . handle , self . _mnff . handle ) )
266+ err , = driver.cuEventRecord (event. handle , self. _handle )
267+ raise_if_driver_error( err )
264268 return event
265269
266270 def wait(self , event_or_stream: Union[Event , Stream]):
@@ -281,7 +285,7 @@ def wait(self, event_or_stream: Union[Event, Stream]):
281285 stream = event_or_stream
282286 else :
283287 try :
284- stream = Stream ._init (event_or_stream )
288+ stream = Stream._init(obj = event_or_stream)
285289 except Exception as e:
286290 raise ValueError (
287291 " only an Event, Stream, or object supporting __cuda_stream__ can be waited,"
@@ -292,7 +296,7 @@ def wait(self, event_or_stream: Union[Event, Stream]):
292296 discard_event = True
293297
294298 # TODO: support flags other than 0?
295- handle_return (driver .cuStreamWaitEvent (self ._mnff . handle , event , 0 ))
299+ handle_return(driver.cuStreamWaitEvent(self ._handle , event, 0 ))
296300 if discard_event:
297301 handle_return(driver.cuEventDestroy(event))
298302
@@ -308,21 +312,22 @@ def device(self) -> Device:
308312
309313 """
310314 from cuda.core.experimental._device import Device # avoid circular import
315+ self._get_device_and_context()
316+ return Device(self._device_id )
311317
318+ cdef _get_device_and_context(self ):
319+ # Get the stream context first
320+ if self ._ctx_handle is None :
321+ err, self ._ctx_handle = driver.cuStreamGetCtx(self ._handle)
322+ raise_if_driver_error(err)
312323 if self ._device_id is None :
313- # Get the stream context first
314- if self ._ctx_handle is None :
315- self ._ctx_handle = handle_return (driver .cuStreamGetCtx (self ._mnff .handle ))
316324 self ._device_id = get_device_from_ctx(self ._ctx_handle)
317- return Device ( self . _device_id )
325+ raise_if_driver_error(err )
318326
319327 @property
320328 def context (self ) -> Context:
321329 """Return the :obj:`~_context.Context` associated with this stream."""
322- if self ._ctx_handle is None :
323- self ._ctx_handle = handle_return (driver .cuStreamGetCtx (self ._mnff .handle ))
324- if self ._device_id is None :
325- self ._device_id = get_device_from_ctx (self ._ctx_handle )
330+ self._get_device_and_context()
326331 return Context._from_ctx(self._ctx_handle , self._device_id )
327332
328333 @staticmethod
0 commit comments