3030import numbers
3131
3232from cpython cimport Py_buffer
33+ from libc.limits cimport INT_MAX
3334from libc.string cimport memcpy
3435
3536from mkl._mkl_service cimport mkl_calloc, mkl_free, mkl_malloc, mkl_realloc
@@ -43,6 +44,37 @@ cdef extern from "stdatomic.h" nogil:
4344 int atomic_load(atomic_int * obj)
4445
4546
47+ cdef extern from * :
48+ """
49+ // Check whether a MKLMemory object may be safely reallocated.
50+ // Mirrors NumPy's PyArray_Resize_int logic.
51+ static int _MKLMemory_MayBeShared(PyObject *op) {
52+ #if PY_VERSION_HEX >= 0x030e00b0
53+ if (PyUnstable_Object_IsUniquelyReferenced(op)) {
54+ return 0; // not shared
55+ }
56+ if (Py_REFCNT(op) == 2) {
57+ return 1; // may be shared
58+ }
59+ return 2; // definitely shared
60+ #else
61+ return (Py_REFCNT(op) > 2) ? 2 : 0;
62+ #endif
63+ }
64+ """
65+ int _MKLMemory_MayBeShared(object obj)
66+
67+
68+ cdef int _check_alignment(Py_ssize_t alignment) except - 1 :
69+ if alignment <= 0 :
70+ raise ValueError (" Alignment of requested allocation must be positive." )
71+ if alignment > < Py_ssize_t> INT_MAX:
72+ raise ValueError (
73+ f" Alignment of requested allocation must not exceed {INT_MAX}."
74+ )
75+ return < int > alignment
76+
77+
4678def _mkl_memory_from_bytes (bytes data , Py_ssize_t alignment ):
4779 cdef Py_ssize_t nbytes = len (data)
4880 cdef MKLMemory mem = MKLMemory(nbytes, alignment = alignment)
@@ -57,28 +89,32 @@ def _mkl_memory_from_bytes(bytes data, Py_ssize_t alignment):
5789
5890
5991cdef class MKLMemory:
92+ """ MKL-backed memory object that exposes Python buffer protocol."""
6093 cdef void * _memory_ptr
61- cdef Py_ssize_t nbytes
62- cdef Py_ssize_t alignment
94+ cdef Py_ssize_t _nbytes
95+ cdef Py_ssize_t _alignment
6396 cdef atomic_int exported_buffers
6497
6598 cdef _cinit_empty(self ):
6699 self ._memory_ptr = NULL
67- self .nbytes = 0
68- self .alignment = 0
100+ self ._nbytes = 0
101+ self ._alignment = 0
69102 atomic_init(& self .exported_buffers, 0 )
70103
71104 cdef _cinit_malloc(self , Py_ssize_t nbytes, Py_ssize_t alignment):
105+ cdef int c_alignment = _check_alignment(alignment)
106+ cdef void * p
107+
72108 self ._cinit_empty()
73109
74110 if (nbytes > 0 ):
75111 with nogil:
76- p = mkl_malloc(nbytes, alignment )
112+ p = mkl_malloc(nbytes, c_alignment )
77113
78114 if (p):
79115 self ._memory_ptr = p
80- self .nbytes = nbytes
81- self .alignment = alignment
116+ self ._nbytes = nbytes
117+ self ._alignment = alignment
82118 else :
83119 raise MemoryError (
84120 " MKL memory allocation failed."
@@ -89,16 +125,19 @@ cdef class MKLMemory:
89125 )
90126
91127 cdef _cinit_calloc(self , Py_ssize_t num, Py_ssize_t size, Py_ssize_t alignment):
128+ cdef int c_alignment = _check_alignment(alignment)
129+ cdef void * p
130+
92131 self ._cinit_empty()
93132
94133 if (num > 0 and size > 0 ):
95134 with nogil:
96- p = mkl_calloc(num, size, alignment )
135+ p = mkl_calloc(num, size, c_alignment )
97136
98137 if (p):
99138 self ._memory_ptr = p
100- self .nbytes = num * size
101- self .alignment = alignment
139+ self ._nbytes = num * size
140+ self ._alignment = alignment
102141 else :
103142 raise MemoryError (
104143 " MKL memory allocation failed."
@@ -110,11 +149,11 @@ cdef class MKLMemory:
110149 )
111150
112151 cdef _cinit_mklmemory(self , object other, Py_ssize_t alignment):
113- other_mem = < MKLMemory> other
152+ cdef MKLMemory other_mem = < MKLMemory> other
114153
115- self ._cinit_malloc(other_mem.nbytes , alignment)
154+ self ._cinit_malloc(other_mem._nbytes , alignment)
116155 with nogil:
117- memcpy(self ._memory_ptr, other_mem._memory_ptr, self .nbytes )
156+ memcpy(self ._memory_ptr, other_mem._memory_ptr, self ._nbytes )
118157
119158 def __cinit__ (self , *args , **kwargs ):
120159 cdef Py_ssize_t alignment
@@ -131,7 +170,7 @@ cdef class MKLMemory:
131170 alignment = kwargs.get(" alignment" , 64 )
132171 self ._cinit_malloc(arg, alignment)
133172 elif isinstance (arg, MKLMemory):
134- alignment = kwargs.get(" alignment" , arg.alignment )
173+ alignment = kwargs.get(" alignment" , ( < MKLMemory > arg)._alignment )
135174 self ._cinit_mklmemory(arg, alignment)
136175 else :
137176 raise TypeError (
@@ -167,11 +206,11 @@ cdef class MKLMemory:
167206 buffer .format = " B"
168207 buffer .internal = NULL
169208 buffer .itemsize = 1
170- buffer .len = self .nbytes
209+ buffer .len = self ._nbytes
171210 buffer .ndim = 1
172211 buffer .obj = self
173212 buffer .readonly = 0
174- buffer .shape = & self .nbytes
213+ buffer .shape = & self ._nbytes
175214 buffer .strides = & buffer .itemsize
176215 buffer .suboffsets = NULL
177216
@@ -181,52 +220,66 @@ cdef class MKLMemory:
181220 atomic_fetch_sub(& self .exported_buffers, 1 )
182221
183222 def realloc (self , Py_ssize_t new_nbytes ):
223+ cdef void * p
224+ cdef int shared
225+
184226 if atomic_load(& self .exported_buffers) > 0 :
185- raise BufferError(" Cannot realloc memory while there are exported buffers." )
227+ raise BufferError(
228+ " Cannot realloc memory while there are exported buffers."
229+ )
230+ shared = _MKLMemory_MayBeShared(self )
231+ if shared == 1 :
232+ raise ValueError (
233+ " Cannot realloc MKLMemory that may be referenced by another "
234+ " object. It is possible that this is a false positive."
235+ )
236+ elif shared == 2 :
237+ raise ValueError (
238+ " Cannot realloc MKLMemory that is referenced by other objects."
239+ )
186240 if new_nbytes <= 0 :
187241 raise ValueError (" New number of bytes must be positive." )
188242
189- cdef void * p
190243 with nogil:
191244 p = mkl_realloc(self ._memory_ptr, new_nbytes)
192245
193246 if not p:
194247 raise MemoryError (" MKL memory reallocation failed." )
195248
196249 self ._memory_ptr = p
197- self .nbytes = new_nbytes
250+ self ._nbytes = new_nbytes
198251
199252 def tobytes (self ):
200253 cdef char * data_ptr = < char * > self ._memory_ptr
201- return data_ptr[:self .nbytes ]
254+ return data_ptr[:self ._nbytes ]
202255
203256 @property
204257 def nbytes (self ):
205- return self .nbytes
258+ return self ._nbytes
206259
207260 @property
208261 def size (self ):
209- return self .nbytes
262+ return self ._nbytes
210263
211264 @property
212265 def alignment (self ):
213- return self .alignment
266+ return self ._alignment
214267
215268 @property
216269 def _pointer (self ):
217270 return < size_t> (self ._memory_ptr)
218271
219272 def __repr__ (self ):
220273 return (
221- f" <MKL memory allocation of {self.nbytes } bytes at "
274+ f" <MKL memory allocation of {self._nbytes } bytes at "
222275 f" {hex(<object>(<size_t>self._memory_ptr))}>"
223276 )
224277
225278 def __len__ (self ):
226- return self .nbytes
279+ return self ._nbytes
227280
228281 def __sizeof__ (self ):
229- return self .nbytes
282+ return self ._nbytes
230283
231284 def __reduce__ (self ):
232- return (_mkl_memory_from_bytes, (self .tobytes(), self .alignment ))
285+ return (_mkl_memory_from_bytes, (self .tobytes(), self ._alignment ))
0 commit comments