@@ -240,91 +240,59 @@ def test_empty(
240240 assert result .flags .f_contiguous # type: ignore[attr-defined]
241241
242242
243- @pytest .mark .parametrize ("dtype" , ["float16" , "float32" , "float64" , "complex64" , "int32" ])
244- @pytest .mark .parametrize ("order" , ["C" , "F" ])
245- def test_all_equal_matches_full_scan (dtype : str , order : str ) -> None :
246- """The blocked scan must agree with a whole-buffer comparison.
247-
248- Buffers are deliberately larger than one block and not a whole number of
249- blocks, and the mismatching element is placed at the very end so an
250- early-exit bug cannot pass by luck.
251- """
252- n = (1 << 14 ) * 3 + 7
253- side = int (np .sqrt (n ))
254- for base in (np .zeros (n , dtype = dtype ), np .zeros ((side , side ), dtype = dtype )):
255- arr = np .asfortranarray (base ) if order == "F" and base .ndim > 1 else base
256- uniform = cpu .NDBuffer .from_numpy_array (arr )
257- differs = arr .copy ()
258- differs .reshape (- 1 )[- 1 ] = 1
259- mixed = cpu .NDBuffer .from_numpy_array (differs )
260- for fill in (0 , 0.0 , 1 ):
261- assert uniform .all_equal (fill ) == cpu .NDBuffer ._compare_all (arr , fill , True )
262- assert mixed .all_equal (fill ) == cpu .NDBuffer ._compare_all (differs , fill , True )
243+ def _all_equal_reference (data : np .ndarray , fill : float ) -> bool :
244+ """`NDBuffer.all_equal` as it was before the fast paths: one full scan,
245+ comparing a zero `fill` bitwise through a void view."""
246+ if np .asarray (fill ).dtype .kind == "f" and fill == 0.0 :
247+ void = f"V{ data .dtype .itemsize } "
248+ expected = np .broadcast_to (np .asarray (fill , data .dtype ), data .shape )
249+ return bool (np .array_equal (data .view (void ), expected .view (void )))
250+ return bool (np .array_equal (data , np .broadcast_to (fill , data .shape ), equal_nan = True ))
251+
252+
253+ _BLOCK = 1 << 14
254+
255+
256+ def _layouts (dtype : str ) -> dict [str , np .ndarray ]:
257+ """Buffers larger than one scan block and not a whole number of blocks."""
258+ side = 211
259+ return {
260+ "1d" : np .zeros (_BLOCK * 3 + 7 , dtype = dtype ),
261+ "C" : np .zeros ((side , side ), dtype = dtype ),
262+ "F" : np .zeros ((side , side ), dtype = dtype , order = "F" ),
263+ # a chunk carved out of a larger array, which is what the write path passes
264+ "strided" : np .zeros ((2 * side , 2 * side ), dtype = dtype )[5 : side + 5 , 7 : side + 7 ],
265+ "leading-1" : np .zeros ((1 , _BLOCK * 2 ), dtype = dtype ),
266+ }
267+
268+
269+ @pytest .mark .parametrize (
270+ "dtype" , ["float16" , "float32" , "float64" , "complex64" , "longdouble" , "int32" ]
271+ )
272+ @pytest .mark .parametrize ("layout" , ["1d" , "C" , "F" , "strided" , "leading-1" ])
273+ @pytest .mark .parametrize ("contents" , [0.0 , - 0.0 , np .nan , 1.0 ])
274+ @pytest .mark .parametrize ("mismatch" , [None , "first" , "last" ])
275+ @pytest .mark .parametrize ("fill" , [0.0 , - 0.0 , np .nan , 1.0 , 0 ])
276+ def test_all_equal (
277+ dtype : str , layout : str , contents : float , mismatch : str | None , fill : float
278+ ) -> None :
279+ """`all_equal` agrees with a full scan, whatever the layout or where a mismatch sits."""
280+ data = _layouts (dtype )[layout ]
281+ if np .dtype (dtype ).kind == "i" and contents != 1.0 :
282+ contents = 0
283+ data [...] = contents
284+ if mismatch is not None :
285+ index = tuple (0 if mismatch == "first" else n - 1 for n in data .shape )
286+ data [index ] = 7
287+ assert cpu .NDBuffer .from_numpy_array (data ).all_equal (fill ) == _all_equal_reference (data , fill )
263288
264289
265290@pytest .mark .parametrize ("dtype" , ["float16" , "float32" , "float64" ])
266291def test_all_equal_distinguishes_negative_zero (dtype : str ) -> None :
267292 """Regression test for #3144: -0.0 is not the same chunk as 0.0."""
268- n = (1 << 14 ) * 2 + 3
269- negative = cpu .NDBuffer .from_numpy_array (np .full (n , - 0.0 , dtype = dtype ))
270- positive = cpu .NDBuffer .from_numpy_array (np .zeros (n , dtype = dtype ))
271- assert positive .all_equal (0.0 )
272- assert not negative .all_equal (0.0 )
273- assert negative .all_equal (- 0.0 )
274- assert not positive .all_equal (- 0.0 )
275- # a single -0.0 in an otherwise +0.0 buffer, past the first block
276- mixed_data = np .zeros (n , dtype = dtype )
277- mixed_data [- 1 ] = np .array (- 0.0 , dtype = dtype )
278- assert not cpu .NDBuffer .from_numpy_array (mixed_data ).all_equal (0.0 )
279-
280-
281- def test_all_equal_non_contiguous () -> None :
282- """A strided buffer cannot be flattened without copying; it must still be correct."""
283- n = (1 << 14 ) * 4
284- base = np .zeros (n * 2 , dtype = "float32" )
285- strided = base [::2 ]
286- assert not strided .flags .c_contiguous
287- assert cpu .NDBuffer .from_numpy_array (strided ).all_equal (0.0 )
288- strided2 = base [::2 ].copy ()
289- strided2 [- 1 ] = 1
290- assert not cpu .NDBuffer .from_numpy_array (strided2 ).all_equal (0.0 )
291-
292-
293- @pytest .mark .parametrize ("fill" , [0.0 , np .nan , 1.0 ])
294- def test_all_equal_strided_chunk_view (fill : float ) -> None :
295- """A chunk carved out of a larger array is strided, which is the shape the
296- write path actually hands this method. Slicing the leading axis has to work
297- for those, not just for contiguous buffers."""
298- whole = np .full ((512 , 512 ), fill , dtype = "float32" )
299- chunk = whole [64 :192 , 64 :192 ]
300- assert not chunk .flags .c_contiguous
301- assert not chunk .flags .f_contiguous
302- assert cpu .NDBuffer .from_numpy_array (chunk ).all_equal (fill )
303- # a single differing element, placed last so an early exit cannot pass by luck
304- whole2 = whole .copy ()
305- whole2 [191 , 191 ] = 12345.0
306- chunk2 = whole2 [64 :192 , 64 :192 ]
307- assert not cpu .NDBuffer .from_numpy_array (chunk2 ).all_equal (fill )
308- # and placed first
309- whole3 = whole .copy ()
310- whole3 [64 , 64 ] = 12345.0
311- assert not cpu .NDBuffer .from_numpy_array (whole3 [64 :192 , 64 :192 ]).all_equal (fill )
312-
313-
314- def test_all_equal_single_leading_row () -> None :
315- """shape[0] == 1 has no leading axis to slice; it must fall back, not break."""
316- row = np .zeros ((1 , (1 << 14 ) * 2 ), dtype = "float32" )
317- assert cpu .NDBuffer .from_numpy_array (row ).all_equal (0.0 )
318- row2 = row .copy ()
319- row2 [0 , - 1 ] = 1
320- assert not cpu .NDBuffer .from_numpy_array (row2 ).all_equal (0.0 )
321-
322-
323- def test_all_equal_nan () -> None :
324- n = (1 << 14 ) * 2 + 1
325- nans = cpu .NDBuffer .from_numpy_array (np .full (n , np .nan , dtype = "float64" ))
326- assert nans .all_equal (np .nan )
327- assert not nans .all_equal (0.0 )
328- one_nan = np .zeros (n , dtype = "float64" )
329- one_nan [- 1 ] = np .nan
330- assert not cpu .NDBuffer .from_numpy_array (one_nan ).all_equal (0.0 )
293+ positive = np .zeros (_BLOCK * 2 + 3 , dtype = dtype )
294+ negative = np .full_like (positive , - 0.0 )
295+ assert cpu .NDBuffer .from_numpy_array (positive ).all_equal (0.0 )
296+ assert not cpu .NDBuffer .from_numpy_array (negative ).all_equal (0.0 )
297+ assert cpu .NDBuffer .from_numpy_array (negative ).all_equal (- 0.0 )
298+ assert not cpu .NDBuffer .from_numpy_array (positive ).all_equal (- 0.0 )
0 commit comments