Skip to content

Commit 3ff5e94

Browse files
committed
cythonize stream + bug fixes
1 parent 7b45954 commit 3ff5e94

6 files changed

Lines changed: 77 additions & 72 deletions

File tree

cuda_core/cuda/core/experimental/_context.pyx

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@ class ContextOptions:
1616
cdef class Context:
1717

1818
cdef:
19-
object _handle
19+
readonly object _handle
2020
int _device_id
2121

2222
def __init__(self, *args, **kwargs):

cuda_core/cuda/core/experimental/_device.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1251,7 +1251,7 @@ def create_stream(self, obj: Optional[IsStreamT] = None, options: StreamOptions
12511251
12521252
"""
12531253
self._check_context_initialized()
1254-
return Stream._init(obj=obj, options=options)
1254+
return Stream._init(obj=obj, options=options, device_id=self._id)
12551255

12561256
def create_event(self, options: Optional[EventOptions] = None) -> Event:
12571257
"""Create an Event object without recording it to a Stream.

cuda_core/cuda/core/experimental/_event.pyx

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -88,19 +88,19 @@ cdef class Event:
8888
raise RuntimeError("Event objects cannot be instantiated directly. Please use Stream APIs (record).")
8989

9090
@classmethod
91-
def _init(cls, device_id: int, ctx_handle: Context, opts=None):
91+
def _init(cls, device_id: int, ctx_handle: Context, options: Optional[EventOptions] = None):
9292
cdef Event self = Event.__new__(Event)
93-
cdef EventOptions options = check_or_create_options(EventOptions, opts, "Event options")
93+
cdef EventOptions opts = check_or_create_options(EventOptions, options, "Event options")
9494
flags = 0x0
9595
self._timing_disabled = False
9696
self._busy_waited = False
97-
if not options.enable_timing:
97+
if not opts.enable_timing:
9898
flags |= driver.CUevent_flags.CU_EVENT_DISABLE_TIMING
9999
self._timing_disabled = True
100-
if options.busy_waited_sync:
100+
if opts.busy_waited_sync:
101101
flags |= driver.CUevent_flags.CU_EVENT_BLOCKING_SYNC
102102
self._busy_waited = True
103-
if options.support_ipc:
103+
if opts.support_ipc:
104104
raise NotImplementedError("WIP: https://github.com/NVIDIA/cuda-python/issues/103")
105105
err, self._handle = driver.cuEventCreate(flags)
106106
raise_if_driver_error(err)

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

Lines changed: 67 additions & 62 deletions
Original file line numberDiff line numberDiff line change
@@ -4,9 +4,13 @@
44

55
from __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+
712
import os
813
import warnings
9-
import weakref
1014
from dataclasses import dataclass
1115
from typing import TYPE_CHECKING, Optional, Protocol, Tuple, Union
1216

@@ -18,16 +22,14 @@
1822
from cuda.core.experimental._graph import GraphBuilder
1923
from cuda.core.experimental._utils.clear_error_support import assert_type
2024
from 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

cuda_core/tests/test_cuda_utils.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -44,7 +44,7 @@ def test_check_driver_error():
4444
num_unexpected = 0
4545
for error in driver.CUresult:
4646
if error == driver.CUresult.CUDA_SUCCESS:
47-
assert cuda_utils._check_driver_error(error) is None
47+
assert cuda_utils._check_driver_error(error) == 0
4848
else:
4949
with pytest.raises(cuda_utils.CUDAError) as e:
5050
cuda_utils._check_driver_error(error)
@@ -63,7 +63,7 @@ def test_check_runtime_error():
6363
num_unexpected = 0
6464
for error in runtime.cudaError_t:
6565
if error == runtime.cudaError_t.cudaSuccess:
66-
assert cuda_utils._check_runtime_error(error) is None
66+
assert cuda_utils._check_runtime_error(error) == 0
6767
else:
6868
with pytest.raises(cuda_utils.CUDAError) as e:
6969
cuda_utils._check_runtime_error(error)

cuda_core/tests/test_stream.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,7 @@ def test_stream_record(init_cuda):
5252

5353
def test_stream_record_invalid_event(init_cuda):
5454
stream = Device().create_stream(options=StreamOptions())
55-
with pytest.raises(TypeError):
55+
with pytest.raises(AttributeError):
5656
stream.record(event="invalid_event")
5757

5858

0 commit comments

Comments
 (0)