Skip to content

Commit 0c75d30

Browse files
committed
align realloc behavior to NumPy
also address issues with undeclared variables and rename MKLMemory class members
1 parent 8cc7b43 commit 0c75d30

2 files changed

Lines changed: 164 additions & 27 deletions

File tree

‎mkl/_mkl_memory.pyx‎

Lines changed: 80 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
import numbers
3131

3232
from cpython cimport Py_buffer
33+
from libc.limits cimport INT_MAX
3334
from libc.string cimport memcpy
3435

3536
from 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+
4678
def _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

5991
cdef 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))

‎mkl/tests/test_mkl_memory.py‎

Lines changed: 84 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,9 @@
2424
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
2525

2626
import sys
27+
import threading
28+
29+
import pytest
2730

2831
import mkl
2932

@@ -137,3 +140,84 @@ def test_pickling_with_alignment():
137140
assert (
138141
mem.alignment == mem_reconstructed.alignment
139142
), "Pickling should preserve alignment"
143+
144+
145+
def test_realloc_exported_buffer():
146+
mem = mkl.MKLMemory(1024)
147+
mv = memoryview(mem)
148+
with pytest.raises(BufferError):
149+
mem.realloc(2048)
150+
del mv
151+
152+
153+
def test_realloc_refcheck_shared():
154+
mem = mkl.MKLMemory(1024)
155+
alias = mem # noqa: F841 — extra reference
156+
with pytest.raises(ValueError, match="referenced by"):
157+
mem.realloc(2048)
158+
del alias
159+
160+
161+
def test_alignment_validation():
162+
with pytest.raises(ValueError, match="positive"):
163+
mkl.MKLMemory(1024, alignment=0)
164+
with pytest.raises(ValueError, match="positive"):
165+
mkl.MKLMemory(1024, alignment=-1)
166+
with pytest.raises(ValueError, match="must not exceed"):
167+
mkl.MKLMemory(1024, alignment=2**40)
168+
169+
170+
def test_concurrent_reads():
171+
mem = mkl.MKLMemory(1024)
172+
mv = memoryview(mem)
173+
for i in range(len(mem)):
174+
mv[i] = i % 256
175+
del mv
176+
177+
errors = []
178+
179+
def reader():
180+
try:
181+
for _ in range(500):
182+
assert len(mem) == 1024
183+
data = mem.tobytes()
184+
assert len(data) == 1024
185+
v = memoryview(mem)
186+
assert v[0] == 0
187+
v.release()
188+
except Exception as e:
189+
errors.append(e)
190+
191+
ts = [threading.Thread(target=reader) for _ in range(4)]
192+
for t in ts:
193+
t.start()
194+
for t in ts:
195+
t.join()
196+
assert not errors, f"Concurrent read errors: {errors}"
197+
198+
199+
def test_concurrent_realloc_refused():
200+
for _ in range(50):
201+
mem = mkl.MKLMemory(64)
202+
barrier = threading.Barrier(2)
203+
results = [None, None]
204+
205+
def worker(idx, size):
206+
barrier.wait()
207+
try:
208+
mem.realloc(size)
209+
results[idx] = "ok"
210+
except ValueError:
211+
results[idx] = "refused"
212+
213+
ts = [
214+
threading.Thread(target=worker, args=(0, 1 << 16)),
215+
threading.Thread(target=worker, args=(1, 1 << 17)),
216+
]
217+
for t in ts:
218+
t.start()
219+
for t in ts:
220+
t.join()
221+
assert results[0] == "refused" and results[1] == "refused", (
222+
f"Expected both threads refused, got {results}"
223+
)

0 commit comments

Comments
 (0)