Skip to content

Commit 67db25e

Browse files
committed
cythonize stream module
1 parent 1976597 commit 67db25e

3 files changed

Lines changed: 72 additions & 30 deletions

File tree

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

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,16 @@
22
#
33
# SPDX-License-Identifier: Apache-2.0
44

5+
from libc.stdint cimport uintptr_t
6+
7+
from cuda.core.experimental._stream cimport _try_to_get_stream_ptr
8+
59
from typing import Union
610

711
from cuda.core.experimental._kernel_arg_handler import ParamHolder
812
from cuda.core.experimental._launch_config import LaunchConfig, _to_native_launch_config
913
from cuda.core.experimental._module import Kernel
10-
from cuda.core.experimental._stream import IsStreamT, Stream, _try_to_get_stream_ptr
14+
from cuda.core.experimental._stream import IsStreamT, Stream
1115
from cuda.core.experimental._utils.clear_error_support import assert_type
1216
from cuda.core.experimental._utils.cuda_utils import (
1317
_reduce_3_tuple,
@@ -60,7 +64,7 @@ def launch(stream: Union[Stream, IsStreamT], config: LaunchConfig, kernel: Kerne
6064
stream_handle = stream.handle
6165
except AttributeError:
6266
try:
63-
stream_handle = _try_to_get_stream_ptr(stream)
67+
stream_handle = driver.CUstream(<uintptr_t>(_try_to_get_stream_ptr(stream)))
6468
except Exception:
6569
raise ValueError(
6670
f"stream must either be a Stream object or support __cuda_stream__ (got {type(stream)})"
Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,9 @@
1+
# SPDX-FileCopyrightText: Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
#
3+
# SPDX-License-Identifier: Apache-2.0
4+
5+
# TODO: how about cuda.bindings < 12.6.2?
6+
from cuda.bindings cimport cydriver
7+
8+
9+
cdef cydriver.CUstream _try_to_get_stream_ptr(obj: IsStreamT) except*

cuda_core/cuda/core/experimental/_stream.pyx

Lines changed: 57 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -4,10 +4,16 @@
44

55
from __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+
712
from cuda.core.experimental._utils.cuda_utils cimport (
813
_check_driver_error as raise_if_driver_error,
914
check_or_create_options,
1015
)
16+
1117
import sys
1218

1319
import 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

9298
cdef 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

Comments
 (0)