codekingpro/portable-devtools
114k
1import itertools
2import os
3import re
4import sys
5import warnings
6import weakref
7
8import pytest
9
10import numpy as np
11import numpy._core._multiarray_umath as ncu
12from numpy.testing import (
13 HAS_REFCOUNT,
14 assert_,
15 assert_allclose,
16 assert_almost_equal,
17 assert_approx_equal,
18 assert_array_almost_equal,
19 assert_array_almost_equal_nulp,
20 assert_array_equal,
21 assert_array_less,
22 assert_array_max_ulp,
23 assert_equal,
24 assert_no_gc_cycles,
25 assert_no_warnings,
26 assert_raises,
27 assert_string_equal,
28 assert_warns,
29 build_err_msg,
30 clear_and_catch_warnings,
31 suppress_warnings,
32 tempdir,
33 temppath,
34)
35
36
37class _GenericTest:
38
39 def _assert_func(self, *args, **kwargs):
40 pass
41
42 def _test_equal(self, a, b):
43 self._assert_func(a, b)
44
45 def _test_not_equal(self, a, b):
46 with assert_raises(AssertionError):
47 self._assert_func(a, b)
48
49 def test_array_rank1_eq(self):
50 """Test two equal array of rank 1 are found equal."""
51 a = np.array([1, 2])
52 b = np.array([1, 2])
53
54 self._test_equal(a, b)
55
56 def test_array_rank1_noteq(self):
57 """Test two different array of rank 1 are found not equal."""
58 a = np.array([1, 2])
59 b = np.array([2, 2])
60
61 self._test_not_equal(a, b)
62
63 def test_array_rank2_eq(self):
64 """Test two equal array of rank 2 are found equal."""
65 a = np.array([[1, 2], [3, 4]])
66 b = np.array([[1, 2], [3, 4]])
67
68 self._test_equal(a, b)
69
70 def test_array_diffshape(self):
71 """Test two arrays with different shapes are found not equal."""
72 a = np.array([1, 2])
73 b = np.array([[1, 2], [1, 2]])
74
75 self._test_not_equal(a, b)
76
77 def test_objarray(self):
78 """Test object arrays."""
79 a = np.array([1, 1], dtype=object)
80 self._test_equal(a, 1)
81
82 def test_array_likes(self):
83 self._test_equal([1, 2, 3], (1, 2, 3))
84
85
86class TestArrayEqual(_GenericTest):
87
88 def _assert_func(self, *args, **kwargs):
89 assert_array_equal(*args, **kwargs)
90
91 def test_generic_rank1(self):
92 """Test rank 1 array for all dtypes."""
93 def foo(t):
94 a = np.empty(2, t)
95 a.fill(1)
96 b = a.copy()
97 c = a.copy()
98 c.fill(0)
99 self._test_equal(a, b)
100 self._test_not_equal(c, b)
101
102 # Test numeric types and object
103 for t in '?bhilqpBHILQPfdgFDG':
104 foo(t)
105
106 # Test strings
107 for t in ['S1', 'U1']:
108 foo(t)
109
110 def test_0_ndim_array(self):
111 x = np.array(473963742225900817127911193656584771)
112 y = np.array(18535119325151578301457182298393896)
113
114 with pytest.raises(AssertionError) as exc_info:
115 self._assert_func(x, y)
116 msg = str(exc_info.value)
117 assert_('Mismatched elements: 1 / 1 (100%)\n'
118 in msg)
119
120 y = x
121 self._assert_func(x, y)
122
123 x = np.array(4395065348745.5643764887869876)
124 y = np.array(0)
125 expected_msg = ('Mismatched elements: 1 / 1 (100%)\n'
126 'Max absolute difference among violations: '
127 '4.39506535e+12\n'
128 'Max relative difference among violations: inf\n')
129 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
130 self._assert_func(x, y)
131
132 x = y
133 self._assert_func(x, y)
134
135 def test_generic_rank3(self):
136 """Test rank 3 array for all dtypes."""
137 def foo(t):
138 a = np.empty((4, 2, 3), t)
139 a.fill(1)
140 b = a.copy()
141 c = a.copy()
142 c.fill(0)
143 self._test_equal(a, b)
144 self._test_not_equal(c, b)
145
146 # Test numeric types and object
147 for t in '?bhilqpBHILQPfdgFDG':
148 foo(t)
149
150 # Test strings
151 for t in ['S1', 'U1']:
152 foo(t)
153
154 def test_nan_array(self):
155 """Test arrays with nan values in them."""
156 a = np.array([1, 2, np.nan])
157 b = np.array([1, 2, np.nan])
158
159 self._test_equal(a, b)
160
161 c = np.array([1, 2, 3])
162 self._test_not_equal(c, b)
163
164 def test_string_arrays(self):
165 """Test two arrays with different shapes are found not equal."""
166 a = np.array(['floupi', 'floupa'])
167 b = np.array(['floupi', 'floupa'])
168
169 self._test_equal(a, b)
170
171 c = np.array(['floupipi', 'floupa'])
172
173 self._test_not_equal(c, b)
174
175 def test_recarrays(self):
176 """Test record arrays."""
177 a = np.empty(2, [('floupi', float), ('floupa', float)])
178 a['floupi'] = [1, 2]
179 a['floupa'] = [1, 2]
180 b = a.copy()
181
182 self._test_equal(a, b)
183
184 c = np.empty(2, [('floupipi', float),
185 ('floupi', float), ('floupa', float)])
186 c['floupipi'] = a['floupi'].copy()
187 c['floupa'] = a['floupa'].copy()
188
189 with pytest.raises(TypeError):
190 self._test_not_equal(c, b)
191
192 def test_masked_nan_inf(self):
193 # Regression test for gh-11121
194 a = np.ma.MaskedArray([3., 4., 6.5], mask=[False, True, False])
195 b = np.array([3., np.nan, 6.5])
196 self._test_equal(a, b)
197 self._test_equal(b, a)
198 a = np.ma.MaskedArray([3., 4., 6.5], mask=[True, False, False])
199 b = np.array([np.inf, 4., 6.5])
200 self._test_equal(a, b)
201 self._test_equal(b, a)
202
203 # Also provides test cases for gh-11121
204 def test_masked_scalar(self):
205 # Test masked scalar vs. plain/masked scalar
206 for a_val, b_val, b_masked in itertools.product(
207 [3., np.nan, np.inf],
208 [3., 4., np.nan, np.inf, -np.inf],
209 [False, True],
210 ):
211 a = np.ma.MaskedArray(a_val, mask=True)
212 b = np.ma.MaskedArray(b_val, mask=True) if b_masked else np.array(b_val)
213 self._test_equal(a, b)
214 self._test_equal(b, a)
215
216 # Test masked scalar vs. plain array
217 for a_val, b_val in itertools.product(
218 [3., np.nan, -np.inf],
219 itertools.product([3., 4., np.nan, np.inf, -np.inf], repeat=2),
220 ):
221 a = np.ma.MaskedArray(a_val, mask=True)
222 b = np.array(b_val)
223 self._test_equal(a, b)
224 self._test_equal(b, a)
225
226 # Test masked scalar vs. masked array
227 for a_val, b_val, b_mask in itertools.product(
228 [3., np.nan, np.inf],
229 itertools.product([3., 4., np.nan, np.inf, -np.inf], repeat=2),
230 itertools.product([False, True], repeat=2),
231 ):
232 a = np.ma.MaskedArray(a_val, mask=True)
233 b = np.ma.MaskedArray(b_val, mask=b_mask)
234 self._test_equal(a, b)
235 self._test_equal(b, a)
236
237 def test_subclass_that_overrides_eq(self):
238 # While we cannot guarantee testing functions will always work for
239 # subclasses, the tests should ideally rely only on subclasses having
240 # comparison operators, not on them being able to store booleans
241 # (which, e.g., astropy Quantity cannot usefully do). See gh-8452.
242 class MyArray(np.ndarray):
243 def __eq__(self, other):
244 return bool(np.equal(self, other).all())
245
246 def __ne__(self, other):
247 return not self == other
248
249 a = np.array([1., 2.]).view(MyArray)
250 b = np.array([2., 3.]).view(MyArray)
251 assert_(type(a == a), bool)
252 assert_(a == a)
253 assert_(a != b)
254 self._test_equal(a, a)
255 self._test_not_equal(a, b)
256 self._test_not_equal(b, a)
257
258 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
259 'Max absolute difference among violations: 1.\n'
260 'Max relative difference among violations: 0.5')
261 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
262 self._test_equal(a, b)
263
264 c = np.array([0., 2.9]).view(MyArray)
265 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
266 'Max absolute difference among violations: 2.\n'
267 'Max relative difference among violations: inf')
268 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
269 self._test_equal(b, c)
270
271 def test_subclass_that_does_not_implement_npall(self):
272 class MyArray(np.ndarray):
273 def __array_function__(self, *args, **kwargs):
274 return NotImplemented
275
276 a = np.array([1., 2.]).view(MyArray)
277 b = np.array([2., 3.]).view(MyArray)
278 with assert_raises(TypeError):
279 np.all(a)
280 self._test_equal(a, a)
281 self._test_not_equal(a, b)
282 self._test_not_equal(b, a)
283
284 def test_suppress_overflow_warnings(self):
285 # Based on issue #18992
286 with pytest.raises(AssertionError):
287 with np.errstate(all="raise"):
288 np.testing.assert_array_equal(
289 np.array([1, 2, 3], np.float32),
290 np.array([1, 1e-40, 3], np.float32))
291
292 def test_array_vs_scalar_is_equal(self):
293 """Test comparing an array with a scalar when all values are equal."""
294 a = np.array([1., 1., 1.])
295 b = 1.
296
297 self._test_equal(a, b)
298
299 def test_array_vs_array_not_equal(self):
300 """Test comparing an array with a scalar when not all values equal."""
301 a = np.array([34986, 545676, 439655, 563766])
302 b = np.array([34986, 545676, 439655, 0])
303
304 expected_msg = ('Mismatched elements: 1 / 4 (25%)\n'
305 'Mismatch at index:\n'
306 ' [3]: 563766 (ACTUAL), 0 (DESIRED)\n'
307 'Max absolute difference among violations: 563766\n'
308 'Max relative difference among violations: inf')
309 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
310 self._assert_func(a, b)
311
312 a = np.array([34986, 545676, 439655.2, 563766])
313 expected_msg = ('Mismatched elements: 2 / 4 (50%)\n'
314 'Mismatch at indices:\n'
315 ' [2]: 439655.2 (ACTUAL), 439655 (DESIRED)\n'
316 ' [3]: 563766.0 (ACTUAL), 0 (DESIRED)\n'
317 'Max absolute difference among violations: '
318 '563766.\n'
319 'Max relative difference among violations: '
320 '4.54902139e-07')
321 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
322 self._assert_func(a, b)
323
324 def test_array_vs_scalar_strict(self):
325 """Test comparing an array with a scalar with strict option."""
326 a = np.array([1., 1., 1.])
327 b = 1.
328
329 with pytest.raises(AssertionError):
330 self._assert_func(a, b, strict=True)
331
332 def test_array_vs_array_strict(self):
333 """Test comparing two arrays with strict option."""
334 a = np.array([1., 1., 1.])
335 b = np.array([1., 1., 1.])
336
337 self._assert_func(a, b, strict=True)
338
339 def test_array_vs_float_array_strict(self):
340 """Test comparing two arrays with strict option."""
341 a = np.array([1, 1, 1])
342 b = np.array([1., 1., 1.])
343
344 with pytest.raises(AssertionError):
345 self._assert_func(a, b, strict=True)
346
347
348class TestBuildErrorMessage:
349
350 def test_build_err_msg_defaults(self):
351 x = np.array([1.00001, 2.00002, 3.00003])
352 y = np.array([1.00002, 2.00003, 3.00004])
353 err_msg = 'There is a mismatch'
354
355 a = build_err_msg([x, y], err_msg)
356 b = ('\nItems are not equal: There is a mismatch\n ACTUAL: array(['
357 '1.00001, 2.00002, 3.00003])\n DESIRED: array([1.00002, '
358 '2.00003, 3.00004])')
359 assert_equal(a, b)
360
361 def test_build_err_msg_no_verbose(self):
362 x = np.array([1.00001, 2.00002, 3.00003])
363 y = np.array([1.00002, 2.00003, 3.00004])
364 err_msg = 'There is a mismatch'
365
366 a = build_err_msg([x, y], err_msg, verbose=False)
367 b = '\nItems are not equal: There is a mismatch'
368 assert_equal(a, b)
369
370 def test_build_err_msg_custom_names(self):
371 x = np.array([1.00001, 2.00002, 3.00003])
372 y = np.array([1.00002, 2.00003, 3.00004])
373 err_msg = 'There is a mismatch'
374
375 a = build_err_msg([x, y], err_msg, names=('FOO', 'BAR'))
376 b = ('\nItems are not equal: There is a mismatch\n FOO: array(['
377 '1.00001, 2.00002, 3.00003])\n BAR: array([1.00002, 2.00003, '
378 '3.00004])')
379 assert_equal(a, b)
380
381 def test_build_err_msg_custom_precision(self):
382 x = np.array([1.000000001, 2.00002, 3.00003])
383 y = np.array([1.000000002, 2.00003, 3.00004])
384 err_msg = 'There is a mismatch'
385
386 a = build_err_msg([x, y], err_msg, precision=10)
387 b = ('\nItems are not equal: There is a mismatch\n ACTUAL: array(['
388 '1.000000001, 2.00002 , 3.00003 ])\n DESIRED: array(['
389 '1.000000002, 2.00003 , 3.00004 ])')
390 assert_equal(a, b)
391
392
393class TestEqual(TestArrayEqual):
394
395 def _assert_func(self, *args, **kwargs):
396 assert_equal(*args, **kwargs)
397
398 def test_nan_items(self):
399 self._assert_func(np.nan, np.nan)
400 self._assert_func([np.nan], [np.nan])
401 self._test_not_equal(np.nan, [np.nan])
402 self._test_not_equal(np.nan, 1)
403
404 def test_inf_items(self):
405 self._assert_func(np.inf, np.inf)
406 self._assert_func([np.inf], [np.inf])
407 self._test_not_equal(np.inf, [np.inf])
408
409 def test_datetime(self):
410 self._test_equal(
411 np.datetime64("2017-01-01", "s"),
412 np.datetime64("2017-01-01", "s")
413 )
414 self._test_equal(
415 np.datetime64("2017-01-01", "s"),
416 np.datetime64("2017-01-01", "m")
417 )
418
419 # gh-10081
420 self._test_not_equal(
421 np.datetime64("2017-01-01", "s"),
422 np.datetime64("2017-01-02", "s")
423 )
424 self._test_not_equal(
425 np.datetime64("2017-01-01", "s"),
426 np.datetime64("2017-01-02", "m")
427 )
428
429 def test_nat_items(self):
430 # not a datetime
431 nadt_no_unit = np.datetime64("NaT")
432 nadt_s = np.datetime64("NaT", "s")
433 nadt_d = np.datetime64("NaT", "ns")
434 # not a timedelta
435 natd_no_unit = np.timedelta64("NaT")
436 natd_s = np.timedelta64("NaT", "s")
437 natd_d = np.timedelta64("NaT", "ns")
438
439 dts = [nadt_no_unit, nadt_s, nadt_d]
440 tds = [natd_no_unit, natd_s, natd_d]
441 for a, b in itertools.product(dts, dts):
442 self._assert_func(a, b)
443 self._assert_func([a], [b])
444 self._test_not_equal([a], b)
445
446 for a, b in itertools.product(tds, tds):
447 self._assert_func(a, b)
448 self._assert_func([a], [b])
449 self._test_not_equal([a], b)
450
451 for a, b in itertools.product(tds, dts):
452 self._test_not_equal(a, b)
453 self._test_not_equal(a, [b])
454 self._test_not_equal([a], [b])
455 self._test_not_equal([a], np.datetime64("2017-01-01", "s"))
456 self._test_not_equal([b], np.datetime64("2017-01-01", "s"))
457 self._test_not_equal([a], np.timedelta64(123, "s"))
458 self._test_not_equal([b], np.timedelta64(123, "s"))
459
460 def test_non_numeric(self):
461 self._assert_func('ab', 'ab')
462 self._test_not_equal('ab', 'abb')
463
464 def test_complex_item(self):
465 self._assert_func(complex(1, 2), complex(1, 2))
466 self._assert_func(complex(1, np.nan), complex(1, np.nan))
467 self._test_not_equal(complex(1, np.nan), complex(1, 2))
468 self._test_not_equal(complex(np.nan, 1), complex(1, np.nan))
469 self._test_not_equal(complex(np.nan, np.inf), complex(np.nan, 2))
470
471 def test_negative_zero(self):
472 self._test_not_equal(ncu.PZERO, ncu.NZERO)
473
474 def test_complex(self):
475 x = np.array([complex(1, 2), complex(1, np.nan)])
476 y = np.array([complex(1, 2), complex(1, 2)])
477 self._assert_func(x, x)
478 self._test_not_equal(x, y)
479
480 def test_object(self):
481 # gh-12942
482 import datetime
483 a = np.array([datetime.datetime(2000, 1, 1),
484 datetime.datetime(2000, 1, 2)])
485 self._test_not_equal(a, a[::-1])
486
487
488class TestArrayAlmostEqual(_GenericTest):
489
490 def _assert_func(self, *args, **kwargs):
491 assert_array_almost_equal(*args, **kwargs)
492
493 def test_closeness(self):
494 # Note that in the course of time we ended up with
495 # `abs(x - y) < 1.5 * 10**(-decimal)`
496 # instead of the previously documented
497 # `abs(x - y) < 0.5 * 10**(-decimal)`
498 # so this check serves to preserve the wrongness.
499
500 # test scalars
501 expected_msg = ('Mismatched elements: 1 / 1 (100%)\n'
502 'Max absolute difference among violations: 1.5\n'
503 'Max relative difference among violations: inf')
504 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
505 self._assert_func(1.5, 0.0, decimal=0)
506
507 # test arrays
508 self._assert_func([1.499999], [0.0], decimal=0)
509
510 expected_msg = ('Mismatched elements: 1 / 1 (100%)\n'
511 'Mismatch at index:\n'
512 ' [0]: 1.5 (ACTUAL), 0.0 (DESIRED)\n'
513 'Max absolute difference among violations: 1.5\n'
514 'Max relative difference among violations: inf')
515 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
516 self._assert_func([1.5], [0.0], decimal=0)
517
518 a = [1.4999999, 0.00003]
519 b = [1.49999991, 0]
520 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
521 'Mismatch at index:\n'
522 ' [1]: 3e-05 (ACTUAL), 0.0 (DESIRED)\n'
523 'Max absolute difference among violations: 3.e-05\n'
524 'Max relative difference among violations: inf')
525 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
526 self._assert_func(a, b, decimal=7)
527
528 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
529 'Mismatch at index:\n'
530 ' [1]: 0.0 (ACTUAL), 3e-05 (DESIRED)\n'
531 'Max absolute difference among violations: 3.e-05\n'
532 'Max relative difference among violations: 1.')
533 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
534 self._assert_func(b, a, decimal=7)
535
536 def test_simple(self):
537 x = np.array([1234.2222])
538 y = np.array([1234.2223])
539
540 self._assert_func(x, y, decimal=3)
541 self._assert_func(x, y, decimal=4)
542
543 expected_msg = ('Mismatched elements: 1 / 1 (100%)\n'
544 'Mismatch at index:\n'
545 ' [0]: 1234.2222 (ACTUAL), 1234.2223 (DESIRED)\n'
546 'Max absolute difference among violations: '
547 '1.e-04\n'
548 'Max relative difference among violations: '
549 '8.10226812e-08')
550 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
551 self._assert_func(x, y, decimal=5)
552
553 def test_array_vs_scalar(self):
554 a = [5498.42354, 849.54345, 0.00]
555 b = 5498.42354
556 expected_msg = ('Mismatched elements: 2 / 3 (66.7%)\n'
557 'Mismatch at indices:\n'
558 ' [1]: 849.54345 (ACTUAL), 5498.42354 (DESIRED)\n'
559 ' [2]: 0.0 (ACTUAL), 5498.42354 (DESIRED)\n'
560 'Max absolute difference among violations: '
561 '5498.42354\n'
562 'Max relative difference among violations: 1.')
563 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
564 self._assert_func(a, b, decimal=9)
565
566 expected_msg = ('Mismatched elements: 2 / 3 (66.7%)\n'
567 'Mismatch at indices:\n'
568 ' [1]: 5498.42354 (ACTUAL), 849.54345 (DESIRED)\n'
569 ' [2]: 5498.42354 (ACTUAL), 0.0 (DESIRED)\n'
570 'Max absolute difference among violations: '
571 '5498.42354\n'
572 'Max relative difference among violations: 5.4722099')
573 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
574 self._assert_func(b, a, decimal=9)
575
576 a = [5498.42354, 0.00]
577 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
578 'Mismatch at index:\n'
579 ' [1]: 5498.42354 (ACTUAL), 0.0 (DESIRED)\n'
580 'Max absolute difference among violations: '
581 '5498.42354\n'
582 'Max relative difference among violations: inf')
583 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
584 self._assert_func(b, a, decimal=7)
585
586 b = 0
587 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
588 'Mismatch at index:\n'
589 ' [0]: 5498.42354 (ACTUAL), 0 (DESIRED)\n'
590 'Max absolute difference among violations: '
591 '5498.42354\n'
592 'Max relative difference among violations: inf')
593 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
594 self._assert_func(a, b, decimal=7)
595
596 def test_nan(self):
597 anan = np.array([np.nan])
598 aone = np.array([1])
599 ainf = np.array([np.inf])
600 self._assert_func(anan, anan)
601 assert_raises(AssertionError,
602 lambda: self._assert_func(anan, aone))
603 assert_raises(AssertionError,
604 lambda: self._assert_func(anan, ainf))
605 assert_raises(AssertionError,
606 lambda: self._assert_func(ainf, anan))
607
608 def test_inf(self):
609 a = np.array([[1., 2.], [3., 4.]])
610 b = a.copy()
611 a[0, 0] = np.inf
612 assert_raises(AssertionError,
613 lambda: self._assert_func(a, b))
614 b[0, 0] = -np.inf
615 assert_raises(AssertionError,
616 lambda: self._assert_func(a, b))
617
618 def test_complex_inf(self):
619 a = np.array([np.inf + 1.j, 2. + 1.j, 3. + 1.j])
620 b = a.copy()
621 self._assert_func(a, b)
622 b[1] = 3. + 1.j
623 expected_msg = ('Mismatched elements: 1 / 3 (33.3%)\n'
624 'Mismatch at index:\n'
625 ' [1]: (2+1j) (ACTUAL), (3+1j) (DESIRED)\n'
626 'Max absolute difference among violations: 1.\n')
627 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
628 self._assert_func(a, b)
629
630 def test_subclass(self):
631 a = np.array([[1., 2.], [3., 4.]])
632 b = np.ma.masked_array([[1., 2.], [0., 4.]],
633 [[False, False], [True, False]])
634 self._assert_func(a, b)
635 self._assert_func(b, a)
636 self._assert_func(b, b)
637
638 # Test fully masked as well (see gh-11123).
639 a = np.ma.MaskedArray(3.5, mask=True)
640 b = np.array([3., 4., 6.5])
641 self._test_equal(a, b)
642 self._test_equal(b, a)
643 a = np.ma.masked
644 b = np.array([3., 4., 6.5])
645 self._test_equal(a, b)
646 self._test_equal(b, a)
647 a = np.ma.MaskedArray([3., 4., 6.5], mask=[True, True, True])
648 b = np.array([1., 2., 3.])
649 self._test_equal(a, b)
650 self._test_equal(b, a)
651 a = np.ma.MaskedArray([3., 4., 6.5], mask=[True, True, True])
652 b = np.array(1.)
653 self._test_equal(a, b)
654 self._test_equal(b, a)
655
656 def test_subclass_2(self):
657 # While we cannot guarantee testing functions will always work for
658 # subclasses, the tests should ideally rely only on subclasses having
659 # comparison operators, not on them being able to store booleans
660 # (which, e.g., astropy Quantity cannot usefully do). See gh-8452.
661 class MyArray(np.ndarray):
662 def __eq__(self, other):
663 return super().__eq__(other).view(np.ndarray)
664
665 def __lt__(self, other):
666 return super().__lt__(other).view(np.ndarray)
667
668 def all(self, *args, **kwargs):
669 return all(self)
670
671 a = np.array([1., 2.]).view(MyArray)
672 self._assert_func(a, a)
673
674 z = np.array([True, True]).view(MyArray)
675 all(z)
676 b = np.array([1., 202]).view(MyArray)
677 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
678 'Mismatch at index:\n'
679 ' [1]: 2.0 (ACTUAL), 202.0 (DESIRED)\n'
680 'Max absolute difference among violations: 200.\n'
681 'Max relative difference among violations: 0.99009')
682 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
683 self._assert_func(a, b)
684
685 def test_subclass_that_cannot_be_bool(self):
686 # While we cannot guarantee testing functions will always work for
687 # subclasses, the tests should ideally rely only on subclasses having
688 # comparison operators, not on them being able to store booleans
689 # (which, e.g., astropy Quantity cannot usefully do). See gh-8452.
690 class MyArray(np.ndarray):
691 def __eq__(self, other):
692 return super().__eq__(other).view(np.ndarray)
693
694 def __lt__(self, other):
695 return super().__lt__(other).view(np.ndarray)
696
697 def all(self, *args, **kwargs):
698 raise NotImplementedError
699
700 a = np.array([1., 2.]).view(MyArray)
701 self._assert_func(a, a)
702
703
704class TestAlmostEqual(_GenericTest):
705
706 def _assert_func(self, *args, **kwargs):
707 assert_almost_equal(*args, **kwargs)
708
709 def test_closeness(self):
710 # Note that in the course of time we ended up with
711 # `abs(x - y) < 1.5 * 10**(-decimal)`
712 # instead of the previously documented
713 # `abs(x - y) < 0.5 * 10**(-decimal)`
714 # so this check serves to preserve the wrongness.
715
716 # test scalars
717 self._assert_func(1.499999, 0.0, decimal=0)
718 assert_raises(AssertionError,
719 lambda: self._assert_func(1.5, 0.0, decimal=0))
720
721 # test arrays
722 self._assert_func([1.499999], [0.0], decimal=0)
723 assert_raises(AssertionError,
724 lambda: self._assert_func([1.5], [0.0], decimal=0))
725
726 def test_nan_item(self):
727 self._assert_func(np.nan, np.nan)
728 assert_raises(AssertionError,
729 lambda: self._assert_func(np.nan, 1))
730 assert_raises(AssertionError,
731 lambda: self._assert_func(np.nan, np.inf))
732 assert_raises(AssertionError,
733 lambda: self._assert_func(np.inf, np.nan))
734
735 def test_inf_item(self):
736 self._assert_func(np.inf, np.inf)
737 self._assert_func(-np.inf, -np.inf)
738 assert_raises(AssertionError,
739 lambda: self._assert_func(np.inf, 1))
740 assert_raises(AssertionError,
741 lambda: self._assert_func(-np.inf, np.inf))
742
743 def test_simple_item(self):
744 self._test_not_equal(1, 2)
745
746 def test_complex_item(self):
747 self._assert_func(complex(1, 2), complex(1, 2))
748 self._assert_func(complex(1, np.nan), complex(1, np.nan))
749 self._assert_func(complex(np.inf, np.nan), complex(np.inf, np.nan))
750 self._test_not_equal(complex(1, np.nan), complex(1, 2))
751 self._test_not_equal(complex(np.nan, 1), complex(1, np.nan))
752 self._test_not_equal(complex(np.nan, np.inf), complex(np.nan, 2))
753
754 def test_complex(self):
755 x = np.array([complex(1, 2), complex(1, np.nan)])
756 z = np.array([complex(1, 2), complex(np.nan, 1)])
757 y = np.array([complex(1, 2), complex(1, 2)])
758 self._assert_func(x, x)
759 self._test_not_equal(x, y)
760 self._test_not_equal(x, z)
761
762 def test_error_message(self):
763 """Check the message is formatted correctly for the decimal value.
764 Also check the message when input includes inf or nan (gh12200)"""
765 x = np.array([1.00000000001, 2.00000000002, 3.00003])
766 y = np.array([1.00000000002, 2.00000000003, 3.00004])
767
768 # Test with a different amount of decimal digits
769 expected_msg = ('Mismatched elements: 3 / 3 (100%)\n'
770 'Mismatch at indices:\n'
771 ' [0]: 1.00000000001 (ACTUAL), 1.00000000002 (DESIRED)\n'
772 ' [1]: 2.00000000002 (ACTUAL), 2.00000000003 (DESIRED)\n'
773 ' [2]: 3.00003 (ACTUAL), 3.00004 (DESIRED)\n'
774 'Max absolute difference among violations: 1.e-05\n'
775 'Max relative difference among violations: '
776 '3.33328889e-06\n'
777 ' ACTUAL: array([1.00000000001, '
778 '2.00000000002, '
779 '3.00003 ])\n'
780 ' DESIRED: array([1.00000000002, 2.00000000003, '
781 '3.00004 ])')
782 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
783 self._assert_func(x, y, decimal=12)
784
785 # With the default value of decimal digits, only the 3rd element
786 # differs. Note that we only check for the formatting of the arrays
787 # themselves.
788 expected_msg = ('Mismatched elements: 1 / 3 (33.3%)\n'
789 'Mismatch at index:\n'
790 ' [2]: 3.00003 (ACTUAL), 3.00004 (DESIRED)\n'
791 'Max absolute difference among violations: 1.e-05\n'
792 'Max relative difference among violations: '
793 '3.33328889e-06\n'
794 ' ACTUAL: array([1. , 2. , 3.00003])\n'
795 ' DESIRED: array([1. , 2. , 3.00004])')
796 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
797 self._assert_func(x, y)
798
799 # Check the error message when input includes inf
800 x = np.array([np.inf, 0])
801 y = np.array([np.inf, 1])
802 expected_msg = ('Mismatched elements: 1 / 2 (50%)\n'
803 'Mismatch at index:\n'
804 ' [1]: 0.0 (ACTUAL), 1.0 (DESIRED)\n'
805 'Max absolute difference among violations: 1.\n'
806 'Max relative difference among violations: 1.\n'
807 ' ACTUAL: array([inf, 0.])\n'
808 ' DESIRED: array([inf, 1.])')
809 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
810 self._assert_func(x, y)
811
812 # Check the error message when dividing by zero
813 x = np.array([1, 2])
814 y = np.array([0, 0])
815 expected_msg = ('Mismatched elements: 2 / 2 (100%)\n'
816 'Mismatch at indices:\n'
817 ' [0]: 1 (ACTUAL), 0 (DESIRED)\n'
818 ' [1]: 2 (ACTUAL), 0 (DESIRED)\n'
819 'Max absolute difference among violations: 2\n'
820 'Max relative difference among violations: inf')
821 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
822 self._assert_func(x, y)
823
824 def test_error_message_2(self):
825 """Check the message is formatted correctly """
826 """when either x or y is a scalar."""
827 x = 2
828 y = np.ones(20)
829 expected_msg = ('Mismatched elements: 20 / 20 (100%)\n'
830 'First 5 mismatches are at indices:\n'
831 ' [0]: 2 (ACTUAL), 1.0 (DESIRED)\n'
832 ' [1]: 2 (ACTUAL), 1.0 (DESIRED)\n'
833 ' [2]: 2 (ACTUAL), 1.0 (DESIRED)\n'
834 ' [3]: 2 (ACTUAL), 1.0 (DESIRED)\n'
835 ' [4]: 2 (ACTUAL), 1.0 (DESIRED)\n'
836 'Max absolute difference among violations: 1.\n'
837 'Max relative difference among violations: 1.')
838 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
839 self._assert_func(x, y)
840
841 y = 2
842 x = np.ones(20)
843 expected_msg = ('Mismatched elements: 20 / 20 (100%)\n'
844 'First 5 mismatches are at indices:\n'
845 ' [0]: 1.0 (ACTUAL), 2 (DESIRED)\n'
846 ' [1]: 1.0 (ACTUAL), 2 (DESIRED)\n'
847 ' [2]: 1.0 (ACTUAL), 2 (DESIRED)\n'
848 ' [3]: 1.0 (ACTUAL), 2 (DESIRED)\n'
849 ' [4]: 1.0 (ACTUAL), 2 (DESIRED)\n'
850 'Max absolute difference among violations: 1.\n'
851 'Max relative difference among violations: 0.5')
852 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
853 self._assert_func(x, y)
854
855 def test_subclass_that_cannot_be_bool(self):
856 # While we cannot guarantee testing functions will always work for
857 # subclasses, the tests should ideally rely only on subclasses having
858 # comparison operators, not on them being able to store booleans
859 # (which, e.g., astropy Quantity cannot usefully do). See gh-8452.
860 class MyArray(np.ndarray):
861 def __eq__(self, other):
862 return super().__eq__(other).view(np.ndarray)
863
864 def __lt__(self, other):
865 return super().__lt__(other).view(np.ndarray)
866
867 def all(self, *args, **kwargs):
868 raise NotImplementedError
869
870 a = np.array([1., 2.]).view(MyArray)
871 self._assert_func(a, a)
872
873
874class TestApproxEqual:
875
876 def _assert_func(self, *args, **kwargs):
877 assert_approx_equal(*args, **kwargs)
878
879 def test_simple_0d_arrays(self):
880 x = np.array(1234.22)
881 y = np.array(1234.23)
882
883 self._assert_func(x, y, significant=5)
884 self._assert_func(x, y, significant=6)
885 assert_raises(AssertionError,
886 lambda: self._assert_func(x, y, significant=7))
887
888 def test_simple_items(self):
889 x = 1234.22
890 y = 1234.23
891
892 self._assert_func(x, y, significant=4)
893 self._assert_func(x, y, significant=5)
894 self._assert_func(x, y, significant=6)
895 assert_raises(AssertionError,
896 lambda: self._assert_func(x, y, significant=7))
897
898 def test_nan_array(self):
899 anan = np.array(np.nan)
900 aone = np.array(1)
901 ainf = np.array(np.inf)
902 self._assert_func(anan, anan)
903 assert_raises(AssertionError, lambda: self._assert_func(anan, aone))
904 assert_raises(AssertionError, lambda: self._assert_func(anan, ainf))
905 assert_raises(AssertionError, lambda: self._assert_func(ainf, anan))
906
907 def test_nan_items(self):
908 anan = np.array(np.nan)
909 aone = np.array(1)
910 ainf = np.array(np.inf)
911 self._assert_func(anan, anan)
912 assert_raises(AssertionError, lambda: self._assert_func(anan, aone))
913 assert_raises(AssertionError, lambda: self._assert_func(anan, ainf))
914 assert_raises(AssertionError, lambda: self._assert_func(ainf, anan))
915
916
917class TestArrayAssertLess:
918
919 def _assert_func(self, *args, **kwargs):
920 assert_array_less(*args, **kwargs)
921
922 def test_simple_arrays(self):
923 x = np.array([1.1, 2.2])
924 y = np.array([1.2, 2.3])
925
926 self._assert_func(x, y)
927 assert_raises(AssertionError, lambda: self._assert_func(y, x))
928
929 y = np.array([1.0, 2.3])
930
931 assert_raises(AssertionError, lambda: self._assert_func(x, y))
932 assert_raises(AssertionError, lambda: self._assert_func(y, x))
933
934 a = np.array([1, 3, 6, 20])
935 b = np.array([2, 4, 6, 8])
936
937 expected_msg = ('Mismatched elements: 2 / 4 (50%)\n'
938 'Mismatch at indices:\n'
939 ' [2]: 6 (x), 6 (y)\n'
940 ' [3]: 20 (x), 8 (y)\n'
941 'Max absolute difference among violations: 12\n'
942 'Max relative difference among violations: 1.5')
943 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
944 self._assert_func(a, b)
945
946 def test_rank2(self):
947 x = np.array([[1.1, 2.2], [3.3, 4.4]])
948 y = np.array([[1.2, 2.3], [3.4, 4.5]])
949
950 self._assert_func(x, y)
951 expected_msg = ('Mismatched elements: 4 / 4 (100%)\n'
952 'Mismatch at indices:\n'
953 ' [0, 0]: 1.2 (x), 1.1 (y)\n'
954 ' [0, 1]: 2.3 (x), 2.2 (y)\n'
955 ' [1, 0]: 3.4 (x), 3.3 (y)\n'
956 ' [1, 1]: 4.5 (x), 4.4 (y)\n'
957 'Max absolute difference among violations: 0.1\n'
958 'Max relative difference among violations: 0.09090909')
959 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
960 self._assert_func(y, x)
961
962 y = np.array([[1.0, 2.3], [3.4, 4.5]])
963 assert_raises(AssertionError, lambda: self._assert_func(x, y))
964 assert_raises(AssertionError, lambda: self._assert_func(y, x))
965
966 def test_rank3(self):
967 x = np.ones(shape=(2, 2, 2))
968 y = np.ones(shape=(2, 2, 2)) + 1
969
970 self._assert_func(x, y)
971 assert_raises(AssertionError, lambda: self._assert_func(y, x))
972
973 y[0, 0, 0] = 0
974 expected_msg = ('Mismatched elements: 1 / 8 (12.5%)\n'
975 'Mismatch at index:\n'
976 ' [0, 0, 0]: 1.0 (x), 0.0 (y)\n'
977 'Max absolute difference among violations: 1.\n'
978 'Max relative difference among violations: inf')
979 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
980 self._assert_func(x, y)
981
982 assert_raises(AssertionError, lambda: self._assert_func(y, x))
983
984 def test_simple_items(self):
985 x = 1.1
986 y = 2.2
987
988 self._assert_func(x, y)
989 expected_msg = ('Mismatched elements: 1 / 1 (100%)\n'
990 'Max absolute difference among violations: 1.1\n'
991 'Max relative difference among violations: 1.')
992 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
993 self._assert_func(y, x)
994
995 y = np.array([2.2, 3.3])
996
997 self._assert_func(x, y)
998 assert_raises(AssertionError, lambda: self._assert_func(y, x))
999
1000 y = np.array([1.0, 3.3])
1001
1002 assert_raises(AssertionError, lambda: self._assert_func(x, y))
1003
1004 def test_simple_items_and_array(self):
1005 x = np.array([[621.345454, 390.5436, 43.54657, 626.4535],
1006 [54.54, 627.3399, 13., 405.5435],
1007 [543.545, 8.34, 91.543, 333.3]])
1008 y = 627.34
1009 self._assert_func(x, y)
1010
1011 y = 8.339999
1012 self._assert_func(y, x)
1013
1014 x = np.array([[3.4536, 2390.5436, 435.54657, 324525.4535],
1015 [5449.54, 999090.54, 130303.54, 405.5435],
1016 [543.545, 8.34, 91.543, 999090.53999]])
1017 y = 999090.54
1018
1019 expected_msg = ('Mismatched elements: 1 / 12 (8.33%)\n'
1020 'Mismatch at index:\n'
1021 ' [1, 1]: 999090.54 (x), 999090.54 (y)\n'
1022 'Max absolute difference among violations: 0.\n'
1023 'Max relative difference among violations: 0.')
1024 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
1025 self._assert_func(x, y)
1026
1027 expected_msg = ('Mismatched elements: 12 / 12 (100%)\n'
1028 'First 5 mismatches are at indices:\n'
1029 ' [0, 0]: 999090.54 (x), 3.4536 (y)\n'
1030 ' [0, 1]: 999090.54 (x), 2390.5436 (y)\n'
1031 ' [0, 2]: 999090.54 (x), 435.54657 (y)\n'
1032 ' [0, 3]: 999090.54 (x), 324525.4535 (y)\n'
1033 ' [1, 0]: 999090.54 (x), 5449.54 (y)\n'
1034 'Max absolute difference among violations: '
1035 '999087.0864\n'
1036 'Max relative difference among violations: '
1037 '289288.5934676')
1038 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
1039 self._assert_func(y, x)
1040
1041 def test_zeroes(self):
1042 x = np.array([546456., 0, 15.455])
1043 y = np.array(87654.)
1044
1045 expected_msg = ('Mismatched elements: 1 / 3 (33.3%)\n'
1046 'Mismatch at index:\n'
1047 ' [0]: 546456.0 (x), 87654.0 (y)\n'
1048 'Max absolute difference among violations: 458802.\n'
1049 'Max relative difference among violations: 5.23423917')
1050 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
1051 self._assert_func(x, y)
1052
1053 expected_msg = ('Mismatched elements: 2 / 3 (66.7%)\n'
1054 'Mismatch at indices:\n'
1055 ' [1]: 87654.0 (x), 0.0 (y)\n'
1056 ' [2]: 87654.0 (x), 15.455 (y)\n'
1057 'Max absolute difference among violations: 87654.\n'
1058 'Max relative difference among violations: '
1059 '5670.5626011')
1060 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
1061 self._assert_func(y, x)
1062
1063 y = 0
1064
1065 expected_msg = ('Mismatched elements: 3 / 3 (100%)\n'
1066 'Mismatch at indices:\n'
1067 ' [0]: 546456.0 (x), 0 (y)\n'
1068 ' [1]: 0.0 (x), 0 (y)\n'
1069 ' [2]: 15.455 (x), 0 (y)\n'
1070 'Max absolute difference among violations: 546456.\n'
1071 'Max relative difference among violations: inf')
1072 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
1073 self._assert_func(x, y)
1074
1075 expected_msg = ('Mismatched elements: 1 / 3 (33.3%)\n'
1076 'Mismatch at index:\n'
1077 ' [1]: 0 (x), 0.0 (y)\n'
1078 'Max absolute difference among violations: 0.\n'
1079 'Max relative difference among violations: inf')
1080 with pytest.raises(AssertionError, match=re.escape(expected_msg)):
1081 self._assert_func(y, x)
1082
1083 def test_nan_noncompare(self):
1084 anan = np.array(np.nan)
1085 aone = np.array(1)
1086 ainf = np.array(np.inf)
1087 self._assert_func(anan, anan)
1088 assert_raises(AssertionError, lambda: self._assert_func(aone, anan))
1089 assert_raises(AssertionError, lambda: self._assert_func(anan, aone))
1090 assert_raises(AssertionError, lambda: self._assert_func(anan, ainf))
1091 assert_raises(AssertionError, lambda: self._assert_func(ainf, anan))
1092
1093 def test_nan_noncompare_array(self):
1094 x = np.array([1.1, 2.2, 3.3])
1095 anan = np.array(np.nan)
1096
1097 assert_raises(AssertionError, lambda: self._assert_func(x, anan))
1098 assert_raises(AssertionError, lambda: self._assert_func(anan, x))
1099
1100 x = np.array([1.1, 2.2, np.nan])
1101
1102 assert_raises(AssertionError, lambda: self._assert_func(x, anan))
1103 assert_raises(AssertionError, lambda: self._assert_func(anan, x))
1104
1105 y = np.array([1.0, 2.0, np.nan])
1106
1107 self._assert_func(y, x)
1108 assert_raises(AssertionError, lambda: self._assert_func(x, y))
1109
1110 def test_inf_compare(self):
1111 aone = np.array(1)
1112 ainf = np.array(np.inf)
1113
1114 self._assert_func(aone, ainf)
1115 self._assert_func(-ainf, aone)
1116 self._assert_func(-ainf, ainf)
1117 assert_raises(AssertionError, lambda: self._assert_func(ainf, aone))
1118 assert_raises(AssertionError, lambda: self._assert_func(aone, -ainf))
1119 assert_raises(AssertionError, lambda: self._assert_func(ainf, ainf))
1120 assert_raises(AssertionError, lambda: self._assert_func(ainf, -ainf))
1121 assert_raises(AssertionError, lambda: self._assert_func(-ainf, -ainf))
1122
1123 def test_inf_compare_array(self):
1124 x = np.array([1.1, 2.2, np.inf])
1125 ainf = np.array(np.inf)
1126
1127 assert_raises(AssertionError, lambda: self._assert_func(x, ainf))
1128 assert_raises(AssertionError, lambda: self._assert_func(ainf, x))
1129 assert_raises(AssertionError, lambda: self._assert_func(x, -ainf))
1130 assert_raises(AssertionError, lambda: self._assert_func(-x, -ainf))
1131 assert_raises(AssertionError, lambda: self._assert_func(-ainf, -x))
1132 self._assert_func(-ainf, x)
1133
1134 def test_strict(self):
1135 """Test the behavior of the `strict` option."""
1136 x = np.zeros(3)
1137 y = np.ones(())
1138 self._assert_func(x, y)
1139 with pytest.raises(AssertionError):
1140 self._assert_func(x, y, strict=True)
1141 y = np.broadcast_to(y, x.shape)
1142 self._assert_func(x, y)
1143 with pytest.raises(AssertionError):
1144 self._assert_func(x, y.astype(np.float32), strict=True)
1145
1146@pytest.mark.filterwarnings(
1147 "ignore:.*NumPy warning suppression and assertion utilities are deprecated"
1148 ".*:DeprecationWarning")
1149@pytest.mark.thread_unsafe(reason="checks global module & deprecated warnings")
1150class TestWarns:
1151
1152 def test_warn(self):
1153 def f():
1154 warnings.warn("yo")
1155 return 3
1156
1157 before_filters = sys.modules['warnings'].filters[:]
1158 assert_equal(assert_warns(UserWarning, f), 3)
1159 after_filters = sys.modules['warnings'].filters
1160
1161 assert_raises(AssertionError, assert_no_warnings, f)
1162 assert_equal(assert_no_warnings(lambda x: x, 1), 1)
1163
1164 # Check that the warnings state is unchanged
1165 assert_equal(before_filters, after_filters,
1166 "assert_warns does not preserver warnings state")
1167
1168 def test_context_manager(self):
1169
1170 before_filters = sys.modules['warnings'].filters[:]
1171 with assert_warns(UserWarning):
1172 warnings.warn("yo")
1173 after_filters = sys.modules['warnings'].filters
1174
1175 def no_warnings():
1176 with assert_no_warnings():
1177 warnings.warn("yo")
1178
1179 assert_raises(AssertionError, no_warnings)
1180 assert_equal(before_filters, after_filters,
1181 "assert_warns does not preserver warnings state")
1182
1183 def test_args(self):
1184 def f(a=0, b=1):
1185 warnings.warn("yo")
1186 return a + b
1187
1188 assert assert_warns(UserWarning, f, b=20) == 20
1189
1190 with pytest.raises(RuntimeError) as exc:
1191 # assert_warns cannot do regexp matching, use pytest.warns
1192 with assert_warns(UserWarning, match="A"):
1193 warnings.warn("B", UserWarning)
1194 assert "assert_warns" in str(exc)
1195 assert "pytest.warns" in str(exc)
1196
1197 with pytest.raises(RuntimeError) as exc:
1198 # assert_warns cannot do regexp matching, use pytest.warns
1199 with assert_warns(UserWarning, wrong="A"):
1200 warnings.warn("B", UserWarning)
