Skip to content

Commit 1bcdf8b

Browse files
authored
Implement structural equality test of memoryview (#1463)
* Implement structural equality test of memoryview * Correctly compare memoryviews containing nans * Refactor BufferBytesEnumerator * Add BytesWarning * Revert BytesWarning, keep test
1 parent 1ca3b36 commit 1bcdf8b

4 files changed

Lines changed: 202 additions & 18 deletions

File tree

Src/IronPython/Runtime/BufferProtocol.cs

Lines changed: 27 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -267,6 +267,9 @@ public static int NumBytes(this IPythonBuffer buffer)
267267
public static BufferBytesEnumerator EnumerateBytes(this IPythonBuffer buffer)
268268
=> new BufferBytesEnumerator(buffer);
269269

270+
public static BufferEnumerator EnumerateItemData(this IPythonBuffer buffer)
271+
=> new BufferEnumerator(buffer, chunkSize: buffer.ItemSize);
272+
270273
/// <summary>
271274
/// Checks if the data in buffer uses a contiguous memory block. If the buffer uses more than one dimension,
272275
/// the data is organized according to the C multi-dimensional array layout.
@@ -325,43 +328,58 @@ public static void CopyTo(this IPythonBuffer buffer, Span<byte> dest) {
325328
}
326329

327330
public ref struct BufferBytesEnumerator {
331+
private readonly BufferEnumerator _enumerator;
332+
333+
public BufferBytesEnumerator(IPythonBuffer buffer)
334+
=> _enumerator = new BufferEnumerator(buffer, chunkSize: 1);
335+
336+
public byte Current => _enumerator.Current[0];
337+
public bool MoveNext() => _enumerator.MoveNext();
338+
public void Dispose() => _enumerator.Dispose();
339+
340+
public BufferBytesEnumerator GetEnumerator() => this;
341+
}
342+
343+
public ref struct BufferEnumerator {
344+
private readonly int _chunksize;
328345
private readonly ReadOnlySpan<byte> _span;
329346
private readonly IEnumerator<int> _offsets;
330347

331-
public BufferBytesEnumerator(IPythonBuffer buffer) {
348+
public BufferEnumerator(IPythonBuffer buffer, int chunkSize) {
332349
if (buffer.SubOffsets != null)
333350
throw new NotImplementedException("buffers with suboffsets are not supported");
334351

352+
_chunksize = chunkSize;
335353
_span = buffer.AsReadOnlySpan();
336-
_offsets = EnumerateDimension(buffer, buffer.Offset, 0).GetEnumerator();
354+
_offsets = EnumerateDimension(buffer, buffer.Offset, chunkSize, 0).GetEnumerator();
337355
}
338356

339-
public byte Current => _span[_offsets.Current];
340-
357+
public ReadOnlySpan<byte> Current => _span.Slice(_offsets.Current, _chunksize);
341358
public bool MoveNext() => _offsets.MoveNext();
359+
public void Dispose() => _offsets.Dispose();
342360

343-
public BufferBytesEnumerator GetEnumerator() => this;
361+
public BufferEnumerator GetEnumerator() => this;
344362

345-
private static IEnumerable<int> EnumerateDimension(IPythonBuffer buffer, int ofs, int dim) {
363+
private static IEnumerable<int> EnumerateDimension(IPythonBuffer buffer, int ofs, int step, int dim) {
346364
IReadOnlyList<int>? shape = buffer.Shape;
347365
IReadOnlyList<int>? strides = buffer.Strides;
348366

349367
if (shape == null || strides == null) {
350368
// simple C-contiguous case
351369
Debug.Assert(buffer.Offset == 0);
352370
int len = buffer.NumBytes();
353-
for (int i = 0; i < len; i++) {
371+
for (int i = 0; i < len; i += step) {
354372
yield return i;
355373
}
356374
} else if (dim >= shape.Count) {
357375
// iterate individual element (scalar)
358-
for (int i = 0; i < buffer.ItemSize; i++) {
376+
for (int i = 0; i < buffer.ItemSize; i += step) {
359377
yield return ofs + i;
360378
}
361379
} else {
362380
for (int i = 0; i < shape[dim]; i++) {
363381
// iterate all bytes from a subdimension
364-
foreach (int j in EnumerateDimension(buffer, ofs, dim + 1)) {
382+
foreach (int j in EnumerateDimension(buffer, ofs, step, dim + 1)) {
365383
yield return j;
366384
}
367385
ofs += strides[dim];

Src/IronPython/Runtime/MemoryView.cs

Lines changed: 49 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -76,8 +76,8 @@ public MemoryView([NotNone] IBufferProtocol @object) {
7676
// for convenience _shape and _strides are never null, even if _numDims == 0 or _flags indicate no _shape or _strides
7777
_shape = _buffer.Shape ?? (_numDims > 0 ? new int[] { _buffer.ItemCount } : Array.Empty<int>());
7878

79-
if (_shape.Count == 0) {
80-
_strides = _shape; // TODO: use a static singleton
79+
if (_numDims == 0) {
80+
_strides = Array.Empty<int>();
8181
_isCContig = true;
8282
} else if (_buffer.Strides != null) {
8383
_strides = _buffer.Strides;
@@ -859,12 +859,55 @@ public int __hash__(CodeContext context) {
859859
return _storedHash.Value;
860860
}
861861

862+
private bool EquivalentShape(MemoryView mv) {
863+
if (_numDims != mv._numDims) return false;
864+
for (int i = 0; i < _numDims; i++) {
865+
if (_shape[i] != mv._shape[i]) return false;
866+
if (_shape[i] == 0) break;
867+
}
868+
return true;
869+
}
870+
862871
public bool __eq__(CodeContext/*!*/ context, [NotNone] MemoryView value) {
863-
if (_buffer == null) {
864-
return value._buffer == null;
872+
if (_buffer == null) return ReferenceEquals(this, value);
873+
if (value._buffer == null) return false;
874+
if (!EquivalentShape(value)) return false;
875+
876+
TypecodeOps.DecomposeTypecode(_format, out char ourByteorder, out char ourTypecode);
877+
// TODO: Support non-native byteorder
878+
879+
// fast tracks if item formats match
880+
if (_format == value._format && !TypecodeOps.IsFloatCode(ourTypecode)) {
881+
if (ReferenceEquals(this, value)) return true;
882+
883+
if (_isCContig && value._isCContig) {
884+
// compare blobs
885+
return ((IPythonBuffer)this).AsReadOnlySpan().SequenceEqual(((IPythonBuffer)value).AsReadOnlySpan());
886+
}
887+
888+
// compare byte by byte
889+
using var ourBytes = this.EnumerateBytes();
890+
using var theirBytes = value.EnumerateBytes();
891+
while (ourBytes.MoveNext() && theirBytes.MoveNext()) {
892+
if (ourBytes.Current != theirBytes.Current) return false;
893+
}
894+
895+
return true;
865896
}
866-
// TODO: comparing flat bytes is oversimplification; besides, no data copyimg
867-
return tobytes().Equals(value.tobytes());
897+
898+
// compare item by item
899+
TypecodeOps.DecomposeTypecode(value._format, out char theirByteorder, out char theirTypecode);
900+
901+
using var us = this.EnumerateItemData();
902+
using var them = value.EnumerateItemData();
903+
while (us.MoveNext() && them.MoveNext()) {
904+
_ = TypecodeOps.TryGetFromBytes(ourTypecode, us.Current, out object? x);
905+
_ = TypecodeOps.TryGetFromBytes(theirTypecode, them.Current, out object? y);
906+
907+
if (!PythonOps.EqualRetBool(x, y)) return false;
908+
}
909+
910+
return true;
868911
}
869912

870913
public bool __eq__(CodeContext/*!*/ context, [NotNone] IBufferProtocol value) => __eq__(context, new MemoryView(value));

Src/IronPython/Runtime/TypecodeOps.cs

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -42,9 +42,11 @@ public static void DecomposeTypecode(string format, out char byteorder, out char
4242
}
4343
}
4444

45-
public static bool IsByteCode(char typecode) {
46-
return typecode == 'B' || typecode == 'b' || typecode == 'c';
47-
}
45+
public static bool IsByteCode(char typecode)
46+
=> typecode == 'B' || typecode == 'b' || typecode == 'c';
47+
48+
public static bool IsFloatCode(char typecode)
49+
=> typecode == 'f' || typecode == 'd';
4850

4951
public static int GetTypecodeWidth(char typecode) {
5052
switch (typecode) {

Tests/test_memoryview.py

Lines changed: 121 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import gc
88
import sys
99
import unittest
10+
import warnings
1011

1112
from iptest import run_test, is_mono, is_cli, is_64
1213

@@ -157,9 +158,129 @@ def test_equality(self):
157158
for x, y in itertools.product((a, mv), repeat=2):
158159
self.assertFalse(x != y, "{!r} {!r}".format(x, y))
159160

161+
def test_equality_structural(self):
160162
# check strided memoryview
161163
self.assertTrue(memoryview(b'axc')[::2] == memoryview(b'ayc')[::2])
162164

165+
# check shape differences
166+
mv = memoryview(b"x")
167+
self.assertEqual(mv.format, "B")
168+
self.assertFalse(mv.cast('B', ()) == mv)
169+
self.assertTrue(mv.cast('B', (1,)) == mv)
170+
self.assertFalse(mv.cast('B', (1,1)) == mv)
171+
172+
b = bytes(range(8))
173+
mv = memoryview(b)
174+
self.assertEqual(mv.format, 'B')
175+
176+
# check different typecodes
177+
mv_b = mv.cast('b')
178+
self.assertTrue(mv_b == mv)
179+
180+
# 'b' to 'B' equivalence does not hold for values out of range
181+
mv_B = memoryview(b'\x80\x81')
182+
mv_b1 = mv_b.cast('b')
183+
self.assertFalse(mv_B == mv_b1)
184+
185+
mv_H = mv.cast('H')
186+
self.assertFalse(mv_H == mv)
187+
mv_h = mv.cast('h')
188+
self.assertTrue(mv_H == mv_h)
189+
190+
mv_i = mv.cast('i')
191+
self.assertFalse(mv_i == mv)
192+
self.assertFalse(mv_i == mv_h)
193+
mv_L = mv.cast('L')
194+
self.assertTrue(mv_i == mv_L)
195+
mv_f = mv.cast('f')
196+
self.assertFalse(mv_i == mv_f)
197+
198+
mv_q = mv.cast('q')
199+
self.assertFalse(mv_q == mv)
200+
self.assertFalse(mv_q == mv_i)
201+
mv_Q = mv.cast('Q')
202+
self.assertTrue(mv_q == mv_Q)
203+
mv_d = mv.cast('d')
204+
self.assertFalse(mv_d == mv_q)
205+
self.assertFalse(mv_d == mv_f)
206+
207+
mv_P = mv.cast('P')
208+
self.assertFalse(mv_P == mv_i)
209+
self.assertFalse(mv_P == mv_L)
210+
self.assertTrue(mv_P == mv_q)
211+
self.assertTrue(mv_P == mv_Q)
212+
213+
# Comparing different formats works if the values are the same
214+
b = bytes(range(8))
215+
mv = memoryview(b)
216+
self.assertEqual(mv.format, 'B')
217+
218+
mv_h = memoryview(array.array('h', [0,1,2,3,4,5,6,7]))
219+
self.assertEqual(mv_h.format, 'h')
220+
self.assertTrue(mv == mv_h)
221+
222+
mv_L = memoryview(array.array('L', [0,1,2,3,4,5,6,7]))
223+
self.assertEqual(mv_L.format, 'L')
224+
self.assertTrue(mv == mv_L)
225+
self.assertTrue(mv_h == mv_L)
226+
227+
mv_d = memoryview(array.array('d', [0,1,2,3,4,5,6,7]))
228+
self.assertEqual(mv_d.format, 'd')
229+
self.assertTrue(mv == mv_d)
230+
self.assertTrue(mv_h == mv_d)
231+
self.assertTrue(mv_L == mv_d)
232+
self.assertTrue(mv_L[::2] == mv_d[::2])
233+
234+
# check released memoryview
235+
ba = bytearray(b)
236+
mv1 = memoryview(b)
237+
mv2 = memoryview(ba)
238+
self.assertTrue(mv1 == mv2)
239+
mv2.release()
240+
self.assertFalse(mv1 == mv2)
241+
mv1.release()
242+
self.assertFalse(mv1 == mv2)
243+
self.assertTrue(mv1 == mv1)
244+
245+
# check nans
246+
z = array.array('f', [float('nan')])
247+
mv = memoryview(z)
248+
self.assertFalse(mv == mv)
249+
mv.release()
250+
self.assertTrue(mv == mv)
251+
252+
@unittest.skipUnless(sys.flags.bytes_warning, "Run Python with the '-b' flag on command line for this test")
253+
def test_equality_warnings(self):
254+
with warnings.catch_warnings(record=True) as ws:
255+
warnings.simplefilter("always")
256+
257+
b = bytes(range(8))
258+
mv = memoryview(b)
259+
self.assertEqual(mv.format, 'B')
260+
261+
mv_b = mv.cast('b')
262+
self.assertTrue(mv_b == mv)
263+
mv_i = mv.cast('i')
264+
self.assertFalse(mv_b == mv_i)
265+
266+
mv_c = mv.cast('c')
267+
if sys.version_info >= (3, 5):
268+
with self.assertWarnsRegex(BytesWarning, r"^Comparison between bytes and int$"):
269+
self.assertFalse(mv_c == mv)
270+
with self.assertWarnsRegex(BytesWarning, r"^Comparison between bytes and int$"):
271+
self.assertFalse(mv_c == mv_b)
272+
with self.assertWarnsRegex(BytesWarning, r"^Comparison between bytes and int$"):
273+
self.assertFalse(mv == mv_c)
274+
with self.assertWarnsRegex(BytesWarning, r"^Comparison between bytes and int$"):
275+
self.assertFalse(mv == mv_c)
276+
else:
277+
self.assertFalse(mv_c == mv)
278+
self.assertFalse(mv_c == mv_b)
279+
self.assertFalse(mv == mv_c)
280+
self.assertFalse(mv == mv_c)
281+
282+
self.assertEqual(len(ws), 0) # no unchecked warnings
283+
163284
def test_overflow(self):
164285
def setitem(m, value):
165286
m[0] = value

0 commit comments

Comments
 (0)