@@ -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
2221cdef 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
153196cdef 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
0 commit comments