Skip to content

Commit 0078990

Browse files
committed
Also cache the cai_data
1 parent d9270a1 commit 0078990

1 file changed

Lines changed: 11 additions & 7 deletions

File tree

cuda_core/cuda/core/experimental/_memoryview.pyx

Lines changed: 11 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -80,11 +80,14 @@ cdef class StridedMemoryView:
8080
bint readonly
8181
object exporting_obj
8282

83+
# If using dlpack, this is a strong reference to the result of
84+
# obj.__dlpack__() so we can lazily create shape and strides from
85+
# it later. If using CAI, this is a reference to the source
86+
# `__cuda_array_interface__` object.
87+
cdef object metadata
88+
8389
# The tensor object if has obj has __dlpack__, otherwise must be NULL
8490
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
8891

8992
# Memoized properties
9093
cdef tuple _shape
@@ -110,7 +113,7 @@ cdef class StridedMemoryView:
110113
self.dl_tensor.ndim
111114
)
112115
else:
113-
self._shape = self.exporting_obj.__cuda_array_interface__["shape"]
116+
self._shape = self.metadata["shape"]
114117
else:
115118
self._shape = ()
116119
return self._shape
@@ -126,7 +129,7 @@ cdef class StridedMemoryView:
126129
self.dl_tensor.ndim
127130
)
128131
else:
129-
strides = self.exporting_obj.__cuda_array_interface__.get("strides")
132+
strides = self.metadata.get("strides")
130133
if strides is not None:
131134
itemsize = self.dtype.itemsize
132135
self._strides = cpython.PyTuple_New(len(strides))
@@ -142,7 +145,7 @@ cdef class StridedMemoryView:
142145
self._dtype = dtype_dlpack_to_numpy(&self.dl_tensor.dtype)
143146
else:
144147
# TODO: this only works for built-in numeric types
145-
self._dtype = numpy.dtype(self.exporting_obj.__cuda_array_interface__["typestr"])
148+
self._dtype = numpy.dtype(self.metadata["typestr"])
146149
return self._dtype
147150

148151
def __repr__(self):
@@ -262,7 +265,7 @@ cdef StridedMemoryView view_as_dlpack(obj, stream_ptr, view=None):
262265

263266
cdef StridedMemoryView buf = StridedMemoryView() if view is None else view
264267
buf.dl_tensor = dl_tensor
265-
buf.dlpack_capsule = capsule
268+
buf.metadata = capsule
266269
buf.ptr = <intptr_t>(dl_tensor.data)
267270
buf.device_id = device_id
268271
buf.is_device_accessible = is_device_accessible
@@ -344,6 +347,7 @@ cpdef StridedMemoryView view_as_cai(obj, stream_ptr, view=None):
344347

345348
cdef StridedMemoryView buf = StridedMemoryView() if view is None else view
346349
buf.exporting_obj = obj
350+
buf.metadata = cai_data
347351
buf.dl_tensor = NULL
348352
buf.ptr, buf.readonly = cai_data["data"]
349353
buf.is_device_accessible = True

0 commit comments

Comments
 (0)