@@ -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