Skip to content

Commit 8c685be

Browse files
committed
Fix #449: Delay construction of Python attributes
1 parent 3443f9a commit 8c685be

2 files changed

Lines changed: 110 additions & 47 deletions

File tree

cuda_core/cuda/core/experimental/_memoryview.pyx

Lines changed: 78 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,6 @@ from cuda.core.experimental._utils cimport cuda_utils
1818
# TODO(leofang): support NumPy structured dtypes
1919

2020

21-
@cython.dataclasses.dataclass
2221
cdef class StridedMemoryView:
2322
"""A dataclass holding metadata of a strided dense array/tensor.
2423
@@ -51,7 +50,7 @@ cdef class StridedMemoryView:
5150
Pointer to the tensor buffer (as a Python `int`).
5251
shape : tuple
5352
Shape of the tensor.
54-
strides : tuple
53+
strides : Optional[tuple]
5554
Strides of the tensor (in **counts**, not bytes).
5655
dtype: numpy.dtype
5756
Data type of the tensor.
@@ -70,19 +69,22 @@ cdef class StridedMemoryView:
7069
obj : Any
7170
Any objects that supports either DLPack (up to v1.0) or CUDA Array
7271
Interface (v3).
73-
stream_ptr: int
72+
stream_ptr: Optional[int]
7473
The pointer address (as Python `int`) to the **consumer** stream.
7574
Stream ordering will be properly established unless ``-1`` is passed.
7675
"""
77-
# TODO: switch to use Cython's cdef typing?
78-
ptr: int = None
79-
shape: tuple = None
80-
strides: tuple = None # in counts, not bytes
81-
dtype: numpy.dtype = None
82-
device_id: int = None # -1 for CPU
83-
is_device_accessible: bool = None
84-
readonly: bool = None
85-
exporting_obj: Any = None
76+
cdef readonly:
77+
intptr_t ptr
78+
int device_id
79+
bint is_device_accessible
80+
bint readonly
81+
object exporting_obj
82+
83+
# The tensor object if has obj has __dlpack__, otherwise must be NULL
84+
cdef DLTensor *dl_tensor
85+
# A strong reference to the result of obj.__dlpack__() so we
86+
# can lazily create shape and strides from it later
87+
cdef object dlpack_capsule
8688

8789
def __init__(self, obj=None, stream_ptr=None):
8890
if obj is not None:
@@ -92,9 +94,50 @@ cdef class StridedMemoryView:
9294
else:
9395
view_as_cai(obj, stream_ptr, self)
9496
else:
95-
# default construct
9697
pass
9798

99+
@property
100+
def shape(self) -> tuple[int]:
101+
if self.exporting_obj is not None:
102+
if self.dl_tensor != NULL:
103+
return cuda_utils.carray_int64_t_to_tuple(
104+
self.dl_tensor.shape,
105+
self.dl_tensor.ndim
106+
)
107+
else:
108+
return self.exporting_obj.__cuda_array_interface__["shape"]
109+
return ()
110+
111+
@property
112+
def strides(self) -> Optional[tuple[int]]:
113+
cdef int itemsize
114+
if self.exporting_obj is not None:
115+
if self.dl_tensor != NULL:
116+
if self.dl_tensor.strides:
117+
return cuda_utils.carray_int64_t_to_tuple(
118+
self.dl_tensor.strides,
119+
self.dl_tensor.ndim
120+
)
121+
else:
122+
strides = self.exporting_obj.__cuda_array_interface__.get("strides")
123+
if strides is not None:
124+
itemsize = self.dtype.itemsize
125+
result = cpython.PyTuple_New(len(strides))
126+
for i in range(len(strides)):
127+
cpython.PyTuple_SET_ITEM(result, i, strides[i] // itemsize)
128+
return result
129+
return None
130+
131+
@property
132+
def dtype(self) -> Optional[numpy.dtype]:
133+
if self.exporting_obj is not None:
134+
if self.dl_tensor != NULL:
135+
return dtype_dlpack_to_numpy(&self.dl_tensor.dtype)
136+
else:
137+
# TODO: this only works for built-in numeric types
138+
return numpy.dtype(self.exporting_obj.__cuda_array_interface__["typestr"])
139+
return None
140+
98141
def __repr__(self):
99142
return (f"StridedMemoryView(ptr={self.ptr},\n"
100143
+ f" shape={self.shape},\n"
@@ -152,7 +195,7 @@ cdef class _StridedMemoryViewProxy:
152195

153196
cdef StridedMemoryView view_as_dlpack(obj, stream_ptr, view=None):
154197
cdef int dldevice, device_id, i
155-
cdef bint is_device_accessible, versioned, is_readonly
198+
cdef bint is_device_accessible, is_readonly
156199
is_device_accessible = False
157200
dldevice, device_id = obj.__dlpack_device__()
158201
if dldevice == _kDLCPU:
@@ -193,7 +236,6 @@ cdef StridedMemoryView view_as_dlpack(obj, stream_ptr, view=None):
193236
capsule, DLPACK_VERSIONED_TENSOR_UNUSED_NAME):
194237
data = cpython.PyCapsule_GetPointer(
195238
capsule, DLPACK_VERSIONED_TENSOR_UNUSED_NAME)
196-
versioned = True
197239
dlm_tensor_ver = <DLManagedTensorVersioned*>data
198240
dl_tensor = &dlm_tensor_ver.dl_tensor
199241
is_readonly = bool((dlm_tensor_ver.flags & DLPACK_FLAG_BITMASK_READ_ONLY) != 0)
@@ -202,32 +244,24 @@ cdef StridedMemoryView view_as_dlpack(obj, stream_ptr, view=None):
202244
capsule, DLPACK_TENSOR_UNUSED_NAME):
203245
data = cpython.PyCapsule_GetPointer(
204246
capsule, DLPACK_TENSOR_UNUSED_NAME)
205-
versioned = False
206247
dlm_tensor = <DLManagedTensor*>data
207248
dl_tensor = &dlm_tensor.dl_tensor
208249
is_readonly = False
209250
used_name = DLPACK_TENSOR_USED_NAME
210251
else:
211252
assert False
212253

254+
cpython.PyCapsule_SetName(capsule, used_name)
255+
213256
cdef StridedMemoryView buf = StridedMemoryView() if view is None else view
257+
buf.dl_tensor = dl_tensor
258+
buf.dlpack_capsule = capsule
214259
buf.ptr = <intptr_t>(dl_tensor.data)
215-
216-
buf.shape = cuda_utils.carray_int64_t_to_tuple(dl_tensor.shape, dl_tensor.ndim)
217-
if dl_tensor.strides:
218-
buf.strides = cuda_utils.carray_int64_t_to_tuple(dl_tensor.strides, dl_tensor.ndim)
219-
else:
220-
# C-order
221-
buf.strides = None
222-
223-
buf.dtype = dtype_dlpack_to_numpy(&dl_tensor.dtype)
224260
buf.device_id = device_id
225261
buf.is_device_accessible = is_device_accessible
226262
buf.readonly = is_readonly
227263
buf.exporting_obj = obj
228264

229-
cpython.PyCapsule_SetName(capsule, used_name)
230-
231265
return buf
232266

233267

@@ -291,7 +325,8 @@ cdef object dtype_dlpack_to_numpy(DLDataType* dtype):
291325
return numpy.dtype(np_dtype)
292326

293327

294-
cdef StridedMemoryView view_as_cai(obj, stream_ptr, view=None):
328+
# Also generate for Python so we can test this code path
329+
cpdef StridedMemoryView view_as_cai(obj, stream_ptr, view=None):
295330
cdef dict cai_data = obj.__cuda_array_interface__
296331
if cai_data["version"] < 3:
297332
raise BufferError("only CUDA Array Interface v3 or above is supported")
@@ -302,33 +337,29 @@ cdef StridedMemoryView view_as_cai(obj, stream_ptr, view=None):
302337

303338
cdef StridedMemoryView buf = StridedMemoryView() if view is None else view
304339
buf.exporting_obj = obj
340+
buf.dl_tensor = NULL
305341
buf.ptr, buf.readonly = cai_data["data"]
306-
buf.shape = cai_data["shape"]
307-
# TODO: this only works for built-in numeric types
308-
buf.dtype = numpy.dtype(cai_data["typestr"])
309-
buf.strides = cai_data.get("strides")
310-
if buf.strides is not None:
311-
# convert to counts
312-
buf.strides = tuple(s // buf.dtype.itemsize for s in buf.strides)
313342
buf.is_device_accessible = True
314343
buf.device_id = handle_return(
315344
driver.cuPointerGetAttribute(
316345
driver.CUpointer_attribute.CU_POINTER_ATTRIBUTE_DEVICE_ORDINAL,
317346
buf.ptr))
318347

319348
cdef intptr_t producer_s, consumer_s
320-
stream = cai_data.get("stream")
321-
if stream is not None:
322-
producer_s = <intptr_t>(stream)
323-
consumer_s = <intptr_t>(stream_ptr)
324-
assert producer_s > 0
325-
# establish stream order
326-
if producer_s != consumer_s:
327-
e = handle_return(driver.cuEventCreate(
328-
driver.CUevent_flags.CU_EVENT_DISABLE_TIMING))
329-
handle_return(driver.cuEventRecord(e, producer_s))
330-
handle_return(driver.cuStreamWaitEvent(consumer_s, e, 0))
331-
handle_return(driver.cuEventDestroy(e))
349+
stream_ptr = int(stream_ptr) if stream_ptr is not None else -1
350+
if stream_ptr != -1:
351+
stream = cai_data.get("stream")
352+
if stream is not None:
353+
producer_s = <intptr_t>(stream)
354+
consumer_s = <intptr_t>(stream_ptr)
355+
assert producer_s > 0
356+
# establish stream order
357+
if producer_s != consumer_s:
358+
e = handle_return(driver.cuEventCreate(
359+
driver.CUevent_flags.CU_EVENT_DISABLE_TIMING))
360+
handle_return(driver.cuEventRecord(e, producer_s))
361+
handle_return(driver.cuStreamWaitEvent(consumer_s, e, 0))
362+
handle_return(driver.cuEventDestroy(e))
332363

333364
return buf
334365

cuda_core/tests/test_utils.py

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515

1616
import cuda.core.experimental
1717
from cuda.core.experimental import Device
18+
from cuda.core.experimental._memoryview import view_as_cai
1819
from cuda.core.experimental.utils import StridedMemoryView, args_viewable_as_strided_memory
1920

2021

@@ -164,3 +165,34 @@ def _check_view(self, view, in_arr, dev):
164165
assert view.is_device_accessible is True
165166
assert view.exporting_obj is in_arr
166167
# can't test view.readonly with CuPy or Numba...
168+
169+
170+
@pytest.mark.skipif(cp is None, reason="CuPy is not installed")
171+
@pytest.mark.parametrize("in_arr,use_stream", (*gpu_array_samples(),))
172+
class TestViewCudaArrayInterfaceGPU:
173+
def test_cuda_array_interface_gpu(self, in_arr, use_stream):
174+
# TODO: use the device fixture?
175+
dev = Device()
176+
dev.set_current()
177+
# This is the consumer stream
178+
s = dev.create_stream() if use_stream else None
179+
180+
# The usual path in `StridedMemoryView` prefers the DLPack interface
181+
# over __cuda_array_interface__, so we call `view_as_cai` directly
182+
# here so we can test the CAI code path.
183+
view = view_as_cai(in_arr, stream_ptr=s.handle if s else -1)
184+
self._check_view(view, in_arr, dev)
185+
186+
def _check_view(self, view, in_arr, dev):
187+
assert isinstance(view, StridedMemoryView)
188+
assert view.ptr == gpu_array_ptr(in_arr)
189+
assert view.shape == in_arr.shape
190+
strides_in_counts = convert_strides_to_counts(in_arr.strides, in_arr.dtype.itemsize)
191+
if in_arr.flags["C_CONTIGUOUS"]:
192+
assert view.strides is None
193+
else:
194+
assert view.strides == strides_in_counts
195+
assert view.dtype == in_arr.dtype
196+
assert view.device_id == dev.device_id
197+
assert view.is_device_accessible is True
198+
assert view.exporting_obj is in_arr

0 commit comments

Comments
 (0)