codekingpro/portable-devtools
115k
1"""
2Utility function to facilitate testing.
3
4"""
5import concurrent.futures
6import contextlib
7import gc
8import importlib.metadata
9import operator
10import os
11import pathlib
12import platform
13import pprint
14import re
15import shutil
16import sys
17import sysconfig
18import threading
19import warnings
20from functools import partial, wraps
21from io import StringIO
22from tempfile import mkdtemp, mkstemp
23from unittest.case import SkipTest
24from warnings import WarningMessage
25
26import numpy as np
27import numpy.linalg._umath_linalg
28from numpy import isfinite, isnan
29from numpy._core import arange, array, array_repr, empty, float32, intp, isnat, ndarray
30
31__all__ = [
32 'assert_equal', 'assert_almost_equal', 'assert_approx_equal',
33 'assert_array_equal', 'assert_array_less', 'assert_string_equal',
34 'assert_array_almost_equal', 'assert_raises', 'build_err_msg',
35 'decorate_methods', 'jiffies', 'memusage', 'print_assert_equal',
36 'rundocs', 'runstring', 'verbose', 'measure',
37 'assert_', 'assert_array_almost_equal_nulp', 'assert_raises_regex',
38 'assert_array_max_ulp', 'assert_warns', 'assert_no_warnings',
39 'assert_allclose', 'IgnoreException', 'clear_and_catch_warnings',
40 'SkipTest', 'KnownFailureException', 'temppath', 'tempdir', 'IS_PYPY',
41 'HAS_REFCOUNT', "IS_WASM", 'suppress_warnings', 'assert_array_compare',
42 'assert_no_gc_cycles', 'break_cycles', 'HAS_LAPACK64', 'IS_PYSTON',
43 'IS_MUSL', 'check_support_sve', 'NOGIL_BUILD',
44 'IS_EDITABLE', 'IS_INSTALLED', 'NUMPY_ROOT', 'run_threaded', 'IS_64BIT',
45 'BLAS_SUPPORTS_FPE',
46 ]
47
48
49class KnownFailureException(Exception):
50 '''Raise this exception to mark a test as a known failing test.'''
51 pass
52
53
54KnownFailureTest = KnownFailureException # backwards compat
55verbose = 0
56
57NUMPY_ROOT = pathlib.Path(np.__file__).parent
58
59try:
60 np_dist = importlib.metadata.distribution('numpy')
61except importlib.metadata.PackageNotFoundError:
62 IS_INSTALLED = IS_EDITABLE = False
63else:
64 IS_INSTALLED = True
65 try:
66 if sys.version_info >= (3, 13):
67 IS_EDITABLE = np_dist.origin.dir_info.editable
68 else:
69 # Backport importlib.metadata.Distribution.origin
70 import json # noqa: E401
71 import types
72 origin = json.loads(
73 np_dist.read_text('direct_url.json') or '{}',
74 object_hook=lambda data: types.SimpleNamespace(**data),
75 )
76 IS_EDITABLE = origin.dir_info.editable
77 except AttributeError:
78 IS_EDITABLE = False
79
80 # spin installs numpy directly via meson, instead of using meson-python, and
81 # runs the module by setting PYTHONPATH. This is problematic because the
82 # resulting installation lacks the Python metadata (.dist-info), and numpy
83 # might already be installed on the environment, causing us to find its
84 # metadata, even though we are not actually loading that package.
85 # Work around this issue by checking if the numpy root matches.
86 if not IS_EDITABLE and np_dist.locate_file('numpy') != NUMPY_ROOT:
87 IS_INSTALLED = False
88
89IS_WASM = platform.machine() in ["wasm32", "wasm64"]
90IS_PYPY = sys.implementation.name == 'pypy'
91IS_PYSTON = hasattr(sys, "pyston_version_info")
92HAS_REFCOUNT = getattr(sys, 'getrefcount', None) is not None and not IS_PYSTON
93BLAS_SUPPORTS_FPE = np._core._multiarray_umath._blas_supports_fpe(None)
94
95HAS_LAPACK64 = numpy.linalg._umath_linalg._ilp64
96
97IS_MUSL = False
98# alternate way is
99# from packaging.tags import sys_tags
100# _tags = list(sys_tags())
101# if 'musllinux' in _tags[0].platform:
102_v = sysconfig.get_config_var('HOST_GNU_TYPE') or ''
103if 'musl' in _v:
104 IS_MUSL = True
105
106NOGIL_BUILD = bool(sysconfig.get_config_var("Py_GIL_DISABLED"))
107IS_64BIT = np.dtype(np.intp).itemsize == 8
108
109def assert_(val, msg=''):
110 """
111 Assert that works in release mode.
112 Accepts callable msg to allow deferring evaluation until failure.
113
114 The Python built-in ``assert`` does not work when executing code in
115 optimized mode (the ``-O`` flag) - no byte-code is generated for it.
116
117 For documentation on usage, refer to the Python documentation.
118
119 """
120 __tracebackhide__ = True # Hide traceback for py.test
121 if not val:
122 try:
123 smsg = msg()
124 except TypeError:
125 smsg = msg
126 raise AssertionError(smsg)
127
128
129if os.name == 'nt':
130 # Code "stolen" from enthought/debug/memusage.py
131 def GetPerformanceAttributes(object, counter, instance=None,
132 inum=-1, format=None, machine=None):
133 # NOTE: Many counters require 2 samples to give accurate results,
134 # including "% Processor Time" (as by definition, at any instant, a
135 # thread's CPU usage is either 0 or 100). To read counters like this,
136 # you should copy this function, but keep the counter open, and call
137 # CollectQueryData() each time you need to know.
138 # See http://msdn.microsoft.com/library/en-us/dnperfmo/html/perfmonpt2.asp
139 # (dead link)
140 # My older explanation for this was that the "AddCounter" process
141 # forced the CPU to 100%, but the above makes more sense :)
142 import win32pdh
143 if format is None:
144 format = win32pdh.PDH_FMT_LONG
145 path = win32pdh.MakeCounterPath((machine, object, instance, None,
146 inum, counter))
147 hq = win32pdh.OpenQuery()
148 try:
149 hc = win32pdh.AddCounter(hq, path)
150 try:
151 win32pdh.CollectQueryData(hq)
152 type, val = win32pdh.GetFormattedCounterValue(hc, format)
153 return val
154 finally:
155 win32pdh.RemoveCounter(hc)
156 finally:
157 win32pdh.CloseQuery(hq)
158
159 def memusage(processName="python", instance=0):
160 # from win32pdhutil, part of the win32all package
161 import win32pdh
162 return GetPerformanceAttributes("Process", "Virtual Bytes",
163 processName, instance,
164 win32pdh.PDH_FMT_LONG, None)
165elif sys.platform[:5] == 'linux':
166
167 def memusage(_proc_pid_stat=None):
168 """
169 Return virtual memory size in bytes of the running python.
170
171 """
172 _proc_pid_stat = _proc_pid_stat or f'/proc/{os.getpid()}/stat'
173 try:
174 with open(_proc_pid_stat) as f:
175 l = f.readline().split(' ')
176 return int(l[22])
177 except Exception:
178 return
179else:
180 def memusage():
181 """
182 Return memory usage of running python. [Not implemented]
183
184 """
185 raise NotImplementedError
186
187
188if sys.platform[:5] == 'linux':
189 def jiffies(_proc_pid_stat=None, _load_time=None):
190 """
191 Return number of jiffies elapsed.
192
193 Return number of jiffies (1/100ths of a second) that this
194 process has been scheduled in user mode. See man 5 proc.
195
196 """
197 _proc_pid_stat = _proc_pid_stat or f'/proc/{os.getpid()}/stat'
198 _load_time = _load_time or []
199 import time
200 if not _load_time:
201 _load_time.append(time.time())
202 try:
203 with open(_proc_pid_stat) as f:
204 l = f.readline().split(' ')
205 return int(l[13])
206 except Exception:
207 return int(100 * (time.time() - _load_time[0]))
208else:
209 # os.getpid is not in all platforms available.
210 # Using time is safe but inaccurate, especially when process
211 # was suspended or sleeping.
212 def jiffies(_load_time=[]):
213 """
214 Return number of jiffies elapsed.
215
216 Return number of jiffies (1/100ths of a second) that this
217 process has been scheduled in user mode. See man 5 proc.
218
219 """
220 import time
221 if not _load_time:
222 _load_time.append(time.time())
223 return int(100 * (time.time() - _load_time[0]))
224
225
226def build_err_msg(arrays, err_msg, header='Items are not equal:',
227 verbose=True, names=('ACTUAL', 'DESIRED'), precision=8):
228 msg = ['\n' + header]
229 err_msg = str(err_msg)
230 if err_msg:
231 if err_msg.find('\n') == -1 and len(err_msg) < 79 - len(header):
232 msg = [msg[0] + ' ' + err_msg]
233 else:
234 msg.append(err_msg)
235 if verbose:
236 for i, a in enumerate(arrays):
237
238 if isinstance(a, ndarray):
239 # precision argument is only needed if the objects are ndarrays
240 r_func = partial(array_repr, precision=precision)
241 else:
242 r_func = repr
243
244 try:
245 r = r_func(a)
246 except Exception as exc:
247 r = f'[repr failed for <{type(a).__name__}>: {exc}]'
248 if r.count('\n') > 3:
249 r = '\n'.join(r.splitlines()[:3])
250 r += '...'
251 msg.append(f' {names[i]}: {r}')
252 return '\n'.join(msg)
253
254
255def assert_equal(actual, desired, err_msg='', verbose=True, *, strict=False):
256 """
257 Raises an AssertionError if two objects are not equal.
258
259 Given two objects (scalars, lists, tuples, dictionaries or numpy arrays),
260 check that all elements of these objects are equal. An exception is raised
261 at the first conflicting values.
262
263 This function handles NaN comparisons as if NaN was a "normal" number.
264 That is, AssertionError is not raised if both objects have NaNs in the same
265 positions. This is in contrast to the IEEE standard on NaNs, which says
266 that NaN compared to anything must return False.
267
268 Parameters
269 ----------
270 actual : array_like
271 The object to check.
272 desired : array_like
273 The expected object.
274 err_msg : str, optional
275 The error message to be printed in case of failure.
276 verbose : bool, optional
277 If True, the conflicting values are appended to the error message.
278 strict : bool, optional
279 If True and either of the `actual` and `desired` arguments is an array,
280 raise an ``AssertionError`` when either the shape or the data type of
281 the arguments does not match. If neither argument is an array, this
282 parameter has no effect.
283
284 .. versionadded:: 2.0.0
285
286 Raises
287 ------
288 AssertionError
289 If actual and desired are not equal.
290
291 See Also
292 --------
293 assert_allclose
294 assert_array_almost_equal_nulp,
295 assert_array_max_ulp,
296
297 Notes
298 -----
299 When one of `actual` and `desired` is a scalar and the other is array_like, the
300 function checks that each element of the array_like is equal to the scalar.
301 Note that empty arrays are therefore considered equal to scalars.
302 This behaviour can be disabled by setting ``strict==True``.
303
304 Examples
305 --------
306 >>> np.testing.assert_equal([4, 5], [4, 6])
307 Traceback (most recent call last):
308 ...
309 AssertionError:
310 Items are not equal:
311 item=1
312 ACTUAL: 5
313 DESIRED: 6
314
315 The following comparison does not raise an exception. There are NaNs
316 in the inputs, but they are in the same positions.
317
318 >>> np.testing.assert_equal(np.array([1.0, 2.0, np.nan]), [1, 2, np.nan])
319
320 As mentioned in the Notes section, `assert_equal` has special
321 handling for scalars when one of the arguments is an array.
322 Here, the test checks that each value in `x` is 3:
323
324 >>> x = np.full((2, 5), fill_value=3)
325 >>> np.testing.assert_equal(x, 3)
326
327 Use `strict` to raise an AssertionError when comparing a scalar with an
328 array of a different shape:
329
330 >>> np.testing.assert_equal(x, 3, strict=True)
331 Traceback (most recent call last):
332 ...
333 AssertionError:
334 Arrays are not equal
335 <BLANKLINE>
336 (shapes (2, 5), () mismatch)
337 ACTUAL: array([[3, 3, 3, 3, 3],
338 [3, 3, 3, 3, 3]])
339 DESIRED: array(3)
340
341 The `strict` parameter also ensures that the array data types match:
342
343 >>> x = np.array([2, 2, 2])
344 >>> y = np.array([2., 2., 2.], dtype=np.float32)
345 >>> np.testing.assert_equal(x, y, strict=True)
346 Traceback (most recent call last):
347 ...
348 AssertionError:
349 Arrays are not equal
350 <BLANKLINE>
351 (dtypes int64, float32 mismatch)
352 ACTUAL: array([2, 2, 2])
353 DESIRED: array([2., 2., 2.], dtype=float32)
354 """
355 __tracebackhide__ = True # Hide traceback for py.test
356 if isinstance(desired, dict):
357 if not isinstance(actual, dict):
358 raise AssertionError(repr(type(actual)))
359 assert_equal(len(actual), len(desired), err_msg, verbose)
360 for k in desired:
361 if k not in actual:
362 raise AssertionError(repr(k))
363 assert_equal(actual[k], desired[k], f'key={k!r}\n{err_msg}',
364 verbose)
365 return
366 if isinstance(desired, (list, tuple)) and isinstance(actual, (list, tuple)):
367 assert_equal(len(actual), len(desired), err_msg, verbose)
368 for k in range(len(desired)):
369 assert_equal(actual[k], desired[k], f'item={k!r}\n{err_msg}',
370 verbose)
371 return
372 from numpy import imag, iscomplexobj, real
373 from numpy._core import isscalar, ndarray, signbit
374 if isinstance(actual, ndarray) or isinstance(desired, ndarray):
375 return assert_array_equal(actual, desired, err_msg, verbose,
376 strict=strict)
377 msg = build_err_msg([actual, desired], err_msg, verbose=verbose)
378
379 # Handle complex numbers: separate into real/imag to handle
380 # nan/inf/negative zero correctly
381 # XXX: catch ValueError for subclasses of ndarray where iscomplex fail
382 try:
383 usecomplex = iscomplexobj(actual) or iscomplexobj(desired)
384 except (ValueError, TypeError):
385 usecomplex = False
386
387 if usecomplex:
388 if iscomplexobj(actual):
389 actualr = real(actual)
390 actuali = imag(actual)
391 else:
392 actualr = actual
393 actuali = 0
394 if iscomplexobj(desired):
395 desiredr = real(desired)
396 desiredi = imag(desired)
397 else:
398 desiredr = desired
399 desiredi = 0
400 try:
401 assert_equal(actualr, desiredr)
402 assert_equal(actuali, desiredi)
403 except AssertionError:
404 raise AssertionError(msg)
405
406 # isscalar test to check cases such as [np.nan] != np.nan
407 if isscalar(desired) != isscalar(actual):
408 raise AssertionError(msg)
409
410 try:
411 isdesnat = isnat(desired)
412 isactnat = isnat(actual)
413 dtypes_match = (np.asarray(desired).dtype.type ==
414 np.asarray(actual).dtype.type)
415 if isdesnat and isactnat:
416 # If both are NaT (and have the same dtype -- datetime or
417 # timedelta) they are considered equal.
418 if dtypes_match:
419 return
420 else:
421 raise AssertionError(msg)
422
423 except (TypeError, ValueError, NotImplementedError):
424 pass
425
426 # Inf/nan/negative zero handling
427 try:
428 isdesnan = isnan(desired)
429 isactnan = isnan(actual)
430 if isdesnan and isactnan:
431 return # both nan, so equal
432
433 # handle signed zero specially for floats
434 array_actual = np.asarray(actual)
435 array_desired = np.asarray(desired)
436 if (array_actual.dtype.char in 'Mm' or
437 array_desired.dtype.char in 'Mm'):
438 # version 1.18
439 # until this version, isnan failed for datetime64 and timedelta64.
440 # Now it succeeds but comparison to scalar with a different type
441 # emits a DeprecationWarning.
442 # Avoid that by skipping the next check
443 raise NotImplementedError('cannot compare to a scalar '
444 'with a different type')
445
446 if desired == 0 and actual == 0:
447 if not signbit(desired) == signbit(actual):
448 raise AssertionError(msg)
449
450 except (TypeError, ValueError, NotImplementedError):
451 pass
452
453 try:
454 # Explicitly use __eq__ for comparison, gh-2552
455 if not (desired == actual):
456 raise AssertionError(msg)
457
458 except (DeprecationWarning, FutureWarning) as e:
459 # this handles the case when the two types are not even comparable
460 if 'elementwise == comparison' in e.args[0]:
461 raise AssertionError(msg)
462 else:
463 raise
464
465
466def print_assert_equal(test_string, actual, desired):
467 """
468 Test if two objects are equal, and print an error message if test fails.
469
470 The test is performed with ``actual == desired``.
471
472 Parameters
473 ----------
474 test_string : str
475 The message supplied to AssertionError.
476 actual : object
477 The object to test for equality against `desired`.
478 desired : object
479 The expected result.
480
481 Examples
482 --------
483 >>> np.testing.print_assert_equal('Test XYZ of func xyz', [0, 1], [0, 1])
484 >>> np.testing.print_assert_equal('Test XYZ of func xyz', [0, 1], [0, 2])
485 Traceback (most recent call last):
486 ...
487 AssertionError: Test XYZ of func xyz failed
488 ACTUAL:
489 [0, 1]
490 DESIRED:
491 [0, 2]
492
493 """
494 __tracebackhide__ = True # Hide traceback for py.test
495 import pprint
496
497 if not (actual == desired):
498 msg = StringIO()
499 msg.write(test_string)
500 msg.write(' failed\nACTUAL: \n')
501 pprint.pprint(actual, msg)
502 msg.write('DESIRED: \n')
503 pprint.pprint(desired, msg)
504 raise AssertionError(msg.getvalue())
505
506
507def assert_almost_equal(actual, desired, decimal=7, err_msg='', verbose=True):
508 """
509 Raises an AssertionError if two items are not equal up to desired
510 precision.
511
512 .. note:: It is recommended to use one of `assert_allclose`,
513 `assert_array_almost_equal_nulp` or `assert_array_max_ulp`
514 instead of this function for more consistent floating point
515 comparisons.
516
517 The test verifies that the elements of `actual` and `desired` satisfy::
518
519 abs(desired-actual) < float64(1.5 * 10**(-decimal))
520
521 That is a looser test than originally documented, but agrees with what the
522 actual implementation in `assert_array_almost_equal` did up to rounding
523 vagaries. An exception is raised at conflicting values. For ndarrays this
524 delegates to assert_array_almost_equal
525
526 Parameters
527 ----------
528 actual : array_like
529 The object to check.
530 desired : array_like
531 The expected object.
532 decimal : int, optional
533 Desired precision, default is 7.
534 err_msg : str, optional
535 The error message to be printed in case of failure.
536 verbose : bool, optional
537 If True, the conflicting values are appended to the error message.
538
539 Raises
540 ------
541 AssertionError
542 If actual and desired are not equal up to specified precision.
543
544 See Also
545 --------
546 assert_allclose: Compare two array_like objects for equality with desired
547 relative and/or absolute precision.
548 assert_array_almost_equal_nulp, assert_array_max_ulp, assert_equal
549
550 Examples
551 --------
552 >>> from numpy.testing import assert_almost_equal
553 >>> assert_almost_equal(2.3333333333333, 2.33333334)
554 >>> assert_almost_equal(2.3333333333333, 2.33333334, decimal=10)
555 Traceback (most recent call last):
556 ...
557 AssertionError:
558 Arrays are not almost equal to 10 decimals
559 ACTUAL: 2.3333333333333
560 DESIRED: 2.33333334
561
562 >>> assert_almost_equal(np.array([1.0,2.3333333333333]),
563 ... np.array([1.0,2.33333334]), decimal=9)
564 Traceback (most recent call last):
565 ...
566 AssertionError:
567 Arrays are not almost equal to 9 decimals
568 <BLANKLINE>
569 Mismatched elements: 1 / 2 (50%)
570 Mismatch at index:
571 [1]: 2.3333333333333 (ACTUAL), 2.33333334 (DESIRED)
572 Max absolute difference among violations: 6.66669964e-09
573 Max relative difference among violations: 2.85715698e-09
574 ACTUAL: array([1. , 2.333333333])
575 DESIRED: array([1. , 2.33333334])
576
577 """
578 __tracebackhide__ = True # Hide traceback for py.test
579 from numpy import imag, iscomplexobj, real
580 from numpy._core import ndarray
581
582 # Handle complex numbers: separate into real/imag to handle
583 # nan/inf/negative zero correctly
584 # XXX: catch ValueError for subclasses of ndarray where iscomplex fail
585 try:
586 usecomplex = iscomplexobj(actual) or iscomplexobj(desired)
587 except ValueError:
588 usecomplex = False
589
590 def _build_err_msg():
591 header = ('Arrays are not almost equal to %d decimals' % decimal)
592 return build_err_msg([actual, desired], err_msg, verbose=verbose,
593 header=header)
594
595 if usecomplex:
596 if iscomplexobj(actual):
597 actualr = real(actual)
598 actuali = imag(actual)
599 else:
600 actualr = actual
601 actuali = 0
602 if iscomplexobj(desired):
603 desiredr = real(desired)
604 desiredi = imag(desired)
605 else:
606 desiredr = desired
607 desiredi = 0
608 try:
609 assert_almost_equal(actualr, desiredr, decimal=decimal)
610 assert_almost_equal(actuali, desiredi, decimal=decimal)
611 except AssertionError:
612 raise AssertionError(_build_err_msg())
613
614 if isinstance(actual, (ndarray, tuple, list)) \
615 or isinstance(desired, (ndarray, tuple, list)):
616 return assert_array_almost_equal(actual, desired, decimal, err_msg)
617 try:
618 # If one of desired/actual is not finite, handle it specially here:
619 # check that both are nan if any is a nan, and test for equality
620 # otherwise
621 if not (isfinite(desired) and isfinite(actual)):
622 if isnan(desired) or isnan(actual):
623 if not (isnan(desired) and isnan(actual)):
624 raise AssertionError(_build_err_msg())
625 elif not desired == actual:
626 raise AssertionError(_build_err_msg())
627 return
628 except (NotImplementedError, TypeError):
629 pass
630 if abs(desired - actual) >= np.float64(1.5 * 10.0**(-decimal)):
631 raise AssertionError(_build_err_msg())
632
633
634def assert_approx_equal(actual, desired, significant=7, err_msg='',
635 verbose=True):
636 """
637 Raises an AssertionError if two items are not equal up to significant
638 digits.
639
640 .. note:: It is recommended to use one of `assert_allclose`,
641 `assert_array_almost_equal_nulp` or `assert_array_max_ulp`
642 instead of this function for more consistent floating point
643 comparisons.
644
645 Given two numbers, check that they are approximately equal.
646 Approximately equal is defined as the number of significant digits
647 that agree.
648
649 Parameters
650 ----------
651 actual : scalar
652 The object to check.
653 desired : scalar
654 The expected object.
655 significant : int, optional
656 Desired precision, default is 7.
657 err_msg : str, optional
658 The error message to be printed in case of failure.
659 verbose : bool, optional
660 If True, the conflicting values are appended to the error message.
661
662 Raises
663 ------
664 AssertionError
665 If actual and desired are not equal up to specified precision.
666
667 See Also
668 --------
669 assert_allclose: Compare two array_like objects for equality with desired
670 relative and/or absolute precision.
671 assert_array_almost_equal_nulp, assert_array_max_ulp, assert_equal
672
673 Examples
674 --------
675 >>> np.testing.assert_approx_equal(0.12345677777777e-20, 0.1234567e-20)
676 >>> np.testing.assert_approx_equal(0.12345670e-20, 0.12345671e-20,
677 ... significant=8)
678 >>> np.testing.assert_approx_equal(0.12345670e-20, 0.12345672e-20,
679 ... significant=8)
680 Traceback (most recent call last):
681 ...
682 AssertionError:
683 Items are not equal to 8 significant digits:
684 ACTUAL: 1.234567e-21
685 DESIRED: 1.2345672e-21
686
687 the evaluated condition that raises the exception is
688
689 >>> abs(0.12345670e-20/1e-21 - 0.12345672e-20/1e-21) >= 10**-(8-1)
690 True
691
692 """
693 __tracebackhide__ = True # Hide traceback for py.test
694 import numpy as np
695
696 (actual, desired) = map(float, (actual, desired))
697 if desired == actual:
698 return
699 # Normalized the numbers to be in range (-10.0,10.0)
700 # scale = float(pow(10,math.floor(math.log10(0.5*(abs(desired)+abs(actual))))))
701 with np.errstate(invalid='ignore'):
702 scale = 0.5 * (np.abs(desired) + np.abs(actual))
703 scale = np.power(10, np.floor(np.log10(scale)))
704 try:
705 sc_desired = desired / scale
706 except ZeroDivisionError:
707 sc_desired = 0.0
708 try:
709 sc_actual = actual / scale
710 except ZeroDivisionError:
711 sc_actual = 0.0
712 msg = build_err_msg(
713 [actual, desired], err_msg,
714 header='Items are not equal to %d significant digits:' % significant,
715 verbose=verbose)
716 try:
717 # If one of desired/actual is not finite, handle it specially here:
718 # check that both are nan if any is a nan, and test for equality
719 # otherwise
720 if not (isfinite(desired) and isfinite(actual)):
721 if isnan(desired) or isnan(actual):
722 if not (isnan(desired) and isnan(actual)):
723 raise AssertionError(msg)
724 elif not desired == actual:
725 raise AssertionError(msg)
726 return
727 except (TypeError, NotImplementedError):
728 pass
729 if np.abs(sc_desired - sc_actual) >= np.power(10., -(significant - 1)):
730 raise AssertionError(msg)
731
732
733def assert_array_compare(comparison, x, y, err_msg='', verbose=True, header='',
734 precision=6, equal_nan=True, equal_inf=True,
735 *, strict=False, names=('ACTUAL', 'DESIRED')):
736 __tracebackhide__ = True # Hide traceback for py.test
737 from numpy._core import all, array2string, errstate, inf, isnan, max, object_
738
739 x = np.asanyarray(x)
740 y = np.asanyarray(y)
741
742 # original array for output formatting
743 ox, oy = x, y
744
745 def isnumber(x):
746 return type(x.dtype)._is_numeric
747
748 def istime(x):
749 return x.dtype.char in "Mm"
750
751 def isvstring(x):
752 return x.dtype.char == "T"
753
754 def robust_any_difference(x, y):
755 # We include work-arounds here to handle three types of slightly
756 # pathological ndarray subclasses:
757 # (1) all() on fully masked arrays returns np.ma.masked, so we use != True
758 # (np.ma.masked != True evaluates as np.ma.masked, which is falsy).
759 # (2) __eq__ on some ndarray subclasses returns Python booleans
760 # instead of element-wise comparisons, so we cast to np.bool() in
761 # that case (or in case __eq__ returns some other value with no
762 # all() method).
763 # (3) subclasses with bare-bones __array_function__ implementations may
764 # not implement np.all(), so favor using the .all() method
765 # We are not committed to supporting cases (2) and (3), but it's nice to
766 # support them if possible.
767 result = x == y
768 if not hasattr(result, "all") or not callable(result.all):
769 result = np.bool(result)
770 return result.all() != True
771
772 def func_assert_same_pos(x, y, func=isnan, hasval='nan'):
773 """Handling nan/inf.
774
775 Combine results of running func on x and y, checking that they are True
776 at the same locations.
777
778 """
779 __tracebackhide__ = True # Hide traceback for py.test
780
781 x_id = func(x)
782 y_id = func(y)
783 if robust_any_difference(x_id, y_id):
784 msg = build_err_msg(
785 [x, y],
786 err_msg + '\n%s location mismatch:'
787 % (hasval), verbose=verbose, header=header,
788 names=names,
789 precision=precision)
790 raise AssertionError(msg)
791 # If there is a scalar, then here we know the array has the same
792 # flag as it everywhere, so we should return the scalar flag.
793 # np.ma.masked is also handled and converted to np.False_ (even if the other
794 # array has nans/infs etc.; that's OK given the handling later of fully-masked
795 # results).
796 if isinstance(x_id, bool) or x_id.ndim == 0:
797 return np.bool(x_id)
798 elif isinstance(y_id, bool) or y_id.ndim == 0:
799 return np.bool(y_id)
800 else:
801 return y_id
802
803 def assert_same_inf_values(x, y, infs_mask):
804 """
805 Verify all inf values match in the two arrays
806 """
807 __tracebackhide__ = True # Hide traceback for py.test
808
809 if not infs_mask.any():
810 return
811 if x.ndim > 0 and y.ndim > 0:
812 x = x[infs_mask]
813 y = y[infs_mask]
814 else:
815 assert infs_mask.all()
816
817 if robust_any_difference(x, y):
818 msg = build_err_msg(
819 [x, y],
820 err_msg + '\ninf values mismatch:',
821 verbose=verbose, header=header,
822 names=names,
823 precision=precision)
824 raise AssertionError(msg)
825
826 try:
827 if strict:
828 cond = x.shape == y.shape and x.dtype == y.dtype
829 else:
830 cond = (x.shape == () or y.shape == ()) or x.shape == y.shape
831 if not cond:
832 if x.shape != y.shape:
833 reason = f'\n(shapes {x.shape}, {y.shape} mismatch)'
834 else:
835 reason = f'\n(dtypes {x.dtype}, {y.dtype} mismatch)'
836 msg = build_err_msg([x, y],
837 err_msg
838 + reason,
839 verbose=verbose, header=header,
840 names=names,
841 precision=precision)
842 raise AssertionError(msg)
843
844 flagged = np.bool(False)
845 if isnumber(x) and isnumber(y):
846 if equal_nan:
847 flagged = func_assert_same_pos(x, y, func=isnan, hasval='nan')
848
849 if equal_inf:
850 # If equal_nan=True, skip comparing nans below for equality if they are
851 # also infs (e.g. inf+nanj) since that would always fail.
852 isinf_func = lambda xy: np.logical_and(np.isinf(xy), np.invert(flagged))
853 infs_mask = func_assert_same_pos(
854 x, y,
855 func=isinf_func,
856 hasval='inf')
857 assert_same_inf_values(x, y, infs_mask)
858 flagged |= infs_mask
859
860 elif istime(x) and istime(y):
861 # If one is datetime64 and the other timedelta64 there is no point
862 if equal_nan and x.dtype.type == y.dtype.type:
863 flagged = func_assert_same_pos(x, y, func=isnat, hasval="NaT")
864
865 elif isvstring(x) and isvstring(y):
866 dt = x.dtype
867 if equal_nan and dt == y.dtype and hasattr(dt, 'na_object'):
868 is_nan = (isinstance(dt.na_object, float) and
869 np.isnan(dt.na_object))
870 bool_errors = 0
871 try:
872 bool(dt.na_object)
873 except TypeError:
874 bool_errors = 1
875 if is_nan or bool_errors:
876 # nan-like NA object
877 flagged = func_assert_same_pos(
878 x, y, func=isnan, hasval=x.dtype.na_object)
879
880 if flagged.ndim > 0:
881 x, y = x[~flagged], y[~flagged]
882 # Only do the comparison if actual values are left
883 if x.size == 0:
884 return
885 elif flagged:
886 # no sense doing comparison if everything is flagged.
887 return
888
889 val = comparison(x, y)
890 invalids = np.logical_not(val)
891
892 if isinstance(val, bool):
893 cond = val
894 reduced = array([val])
895 else:
896 reduced = val.ravel()
897 cond = reduced.all()
898
899 # The below comparison is a hack to ensure that fully masked
900 # results, for which val.ravel().all() returns np.ma.masked,
901 # do not trigger a failure (np.ma.masked != True evaluates as
902 # np.ma.masked, which is falsy).
903 if cond != True:
904 n_mismatch = reduced.size - reduced.sum(dtype=intp)
905 n_elements = flagged.size if flagged.ndim != 0 else reduced.size
906 percent_mismatch = 100 * n_mismatch / n_elements
907 remarks = [f'Mismatched elements: {n_mismatch} / {n_elements} '
908 f'({percent_mismatch:.3g}%)']
909 if invalids.ndim != 0:
910 if flagged.ndim > 0:
911 positions = np.argwhere(np.asarray(~flagged))[invalids]
912 else:
913 positions = np.argwhere(np.asarray(invalids))
914 s = "\n".join(
915 [
916 f" {p.tolist()}: {ox if ox.ndim == 0 else ox[tuple(p)]} "
917 f"({names[0]}), {oy if oy.ndim == 0 else oy[tuple(p)]} "
918 f"({names[1]})"
919 for p in positions[:5]
920 ]
921 )
922 if len(positions) == 1:
923 remarks.append(
924 f"Mismatch at index:\n{s}"
925 )
926 elif len(positions) <= 5:
927 remarks.append(
928 f"Mismatch at indices:\n{s}"
929 )
930 else:
931 remarks.append(
932 f"First 5 mismatches are at indices:\n{s}"
933 )
934
935 with errstate(all='ignore'):
936 # ignore errors for non-numeric types
937 with contextlib.suppress(TypeError):
938 error = abs(x - y)
939 if np.issubdtype(x.dtype, np.unsignedinteger):
940 error2 = abs(y - x)
941 np.minimum(error, error2, out=error)
942
943 reduced_error = error[invalids]
944 max_abs_error = max(reduced_error)
945 if getattr(error, 'dtype', object_) == object_:
946 remarks.append(
947 'Max absolute difference among violations: '
948 + str(max_abs_error))
949 else:
950 remarks.append(
951 'Max absolute difference among violations: '
952 + array2string(max_abs_error))
953
954 # note: this definition of relative error matches that one
955 # used by assert_allclose (found in np.isclose)
956 # Filter values where the divisor would be zero
957 nonzero = np.bool(y != 0)
958 nonzero_and_invalid = np.logical_and(invalids, nonzero)
959
960 if all(~nonzero_and_invalid):
961 max_rel_error = array(inf)
962 else:
963 nonzero_invalid_error = error[nonzero_and_invalid]
964 broadcasted_y = np.broadcast_to(y, error.shape)
965 nonzero_invalid_y = broadcasted_y[nonzero_and_invalid]
966 max_rel_error = max(nonzero_invalid_error
967 / abs(nonzero_invalid_y))
968
969 if getattr(error, 'dtype', object_) == object_:
970 remarks.append(
971 'Max relative difference among violations: '
972 + str(max_rel_error))
973 else:
974 remarks.append(
975 'Max relative difference among violations: '
976 + array2string(max_rel_error))
977 err_msg = str(err_msg)
978 err_msg += '\n' + '\n'.join(remarks)
979 msg = build_err_msg([ox, oy], err_msg,
980 verbose=verbose, header=header,
981 names=names,
982 precision=precision)
983 raise AssertionError(msg)
984 except ValueError:
985 import traceback
986 efmt = traceback.format_exc()
987 header = f'error during assertion:\n\n{efmt}\n\n{header}'
988
989 msg = build_err_msg([x, y], err_msg, verbose=verbose, header=header,
990 names=names, precision=precision)
991 raise ValueError(msg)
992
993
994def assert_array_equal(actual, desired, err_msg='', verbose=True, *,
995 strict=False):
996 """
997 Raises an AssertionError if two array_like objects are not equal.
998
999 Given two array_like objects, check that the shape is equal and all
1000 elements of these objects are equal (but see the Notes for the special
1001 handling of a scalar). An exception is raised at shape mismatch or
1002 conflicting values. In contrast to the standard usage in numpy, NaNs
1003 are compared like numbers, no assertion is raised if both objects have
1004 NaNs in the same positions.
1005
1006 The usual caution for verifying equality with floating point numbers is
1007 advised.
1008
1009 .. note:: When either `actual` or `desired` is already an instance of
1010 `numpy.ndarray` and `desired` is not a ``dict``, the behavior of
1011 ``assert_equal(actual, desired)`` is identical to the behavior of this
1012 function. Otherwise, this function performs `np.asanyarray` on the
1013 inputs before comparison, whereas `assert_equal` defines special
1014 comparison rules for common Python types. For example, only
1015 `assert_equal` can be used to compare nested Python lists. In new code,
1016 consider using only `assert_equal`, explicitly converting either
1017 `actual` or `desired` to arrays if the behavior of `assert_array_equal`
1018 is desired.
1019
1020 Parameters
1021 ----------
1022 actual : array_like
1023 The actual object to check.
1024 desired : array_like
1025 The desired, expected object.
1026 err_msg : str, optional
1027 The error message to be printed in case of failure.
1028 verbose : bool, optional
1029 If True, the conflicting values are appended to the error message.
1030 strict : bool, optional
1031 If True, raise an AssertionError when either the shape or the data
1032 type of the array_like objects does not match. The special
1033 handling for scalars mentioned in the Notes section is disabled.
1034
1035 .. versionadded:: 1.24.0
1036
1037 Raises
1038 ------
1039 AssertionError
1040 If actual and desired objects are not equal.
1041
1042 See Also
1043 --------
1044 assert_allclose: Compare two array_like objects for equality with desired
1045 relative and/or absolute precision.
1046 assert_array_almost_equal_nulp, assert_array_max_ulp, assert_equal
1047
1048 Notes
1049 -----
1050 When one of `actual` and `desired` is a scalar and the other is array_like, the
1051 function checks that each element of the array_like is equal to the scalar.
1052 Note that empty arrays are therefore considered equal to scalars.
1053 This behaviour can be disabled by setting ``strict==True``.
1054
1055 Examples
1056 --------
1057 The first assert does not raise an exception:
1058
1059 >>> np.testing.assert_array_equal([1.0,2.33333,np.nan],
1060 ... [np.exp(0),2.33333, np.nan])
1061
1062 Assert fails with numerical imprecision with floats:
1063
1064 >>> np.testing.assert_array_equal([1.0,np.pi,np.nan],
1065 ... [1, np.sqrt(np.pi)**2, np.nan])
1066 Traceback (most recent call last):
1067 ...
1068 AssertionError:
1069 Arrays are not equal
1070 <BLANKLINE>
1071 Mismatched elements: 1 / 3 (33.3%)
1072 Mismatch at index:
1073 [1]: 3.141592653589793 (ACTUAL), 3.1415926535897927 (DESIRED)
1074 Max absolute difference among violations: 4.4408921e-16
1075 Max relative difference among violations: 1.41357986e-16
1076 ACTUAL: array([1. , 3.141593, nan])
1077 DESIRED: array([1. , 3.141593, nan])
1078
1079 Use `assert_allclose` or one of the nulp (number of floating point values)
1080 functions for these cases instead:
1081
1082 >>> np.testing.assert_allclose([1.0,np.pi,np.nan],
1083 ... [1, np.sqrt(np.pi)**2, np.nan],
1084 ... rtol=1e-10, atol=0)
1085
1086 As mentioned in the Notes section, `assert_array_equal` has special
1087 handling for scalars. Here the test checks that each value in `x` is 3:
1088
1089 >>> x = np.full((2, 5), fill_value=3)
1090 >>> np.testing.assert_array_equal(x, 3)
1091
1092 Use `strict` to raise an AssertionError when comparing a scalar with an
1093 array:
1094
1095 >>> np.testing.assert_array_equal(x, 3, strict=True)
1096 Traceback (most recent call last):
1097 ...
1098 AssertionError:
1099 Arrays are not equal
1100 <BLANKLINE>
1101 (shapes (2, 5), () mismatch)
1102 ACTUAL: array([[3, 3, 3, 3, 3],
1103 [3, 3, 3, 3, 3]])
1104 DESIRED: array(3)
1105
1106 The `strict` parameter also ensures that the array data types match:
1107
1108 >>> x = np.array([2, 2, 2])
1109 >>> y = np.array([2., 2., 2.], dtype=np.float32)
1110 >>> np.testing.assert_array_equal(x, y, strict=True)
1111 Traceback (most recent call last):
1112 ...
1113 AssertionError:
1114 Arrays are not equal
1115 <BLANKLINE>
1116 (dtypes int64, float32 mismatch)
1117 ACTUAL: array([2, 2, 2])
1118 DESIRED: array([2., 2., 2.], dtype=float32)
1119 """
1120 __tracebackhide__ = True # Hide traceback for py.test
1121 assert_array_compare(operator.__eq__, actual, desired, err_msg=err_msg,
1122 verbose=verbose, header='Arrays are not equal',
1123 strict=strict)
1124
1125
1126def assert_array_almost_equal(actual, desired, decimal=6, err_msg='',
1127 verbose=True):
1128 """
1129 Raises an AssertionError if two objects are not equal up to desired
1130 precision.
1131
1132 .. note:: It is recommended to use one of `assert_allclose`,
1133 `assert_array_almost_equal_nulp` or `assert_array_max_ulp`
1134 instead of this function for more consistent floating point
1135 comparisons.
1136
1137 The test verifies identical shapes and that the elements of ``actual`` and
1138 ``desired`` satisfy::
1139
1140 abs(desired-actual) < 1.5 * 10**(-decimal)
1141
1142 That is a looser test than originally documented, but agrees with what the
1143 actual implementation did up to rounding vagaries. An exception is raised
1144 at shape mismatch or conflicting values. In contrast to the standard usage
1145 in numpy, NaNs are compared like numbers, no assertion is raised if both
1146 objects have NaNs in the same positions.
1147
1148 Parameters
1149 ----------
1150 actual : array_like
1151 The actual object to check.
1152 desired : array_like
1153 The desired, expected object.
1154 decimal : int, optional
1155 Desired precision, default is 6.
1156 err_msg : str, optional
1157 The error message to be printed in case of failure.
1158 verbose : bool, optional
1159 If True, the conflicting values are appended to the error message.
1160
1161 Raises
1162 ------
1163 AssertionError
1164 If actual and desired are not equal up to specified precision.
1165
1166 See Also
1167 --------
1168 assert_allclose: Compare two array_like objects for equality with desired
1169 relative and/or absolute precision.
1170 assert_array_almost_equal_nulp, assert_array_max_ulp, assert_equal
1171
1172 Examples
1173 --------
1174 the first assert does not raise an exception
1175
1176 >>> np.testing.assert_array_almost_equal([1.0,2.333,np.nan],
1177 ... [1.0,2.333,np.nan])
1178
1179 >>> np.testing.assert_array_almost_equal([1.0,2.33333,np.nan],
1180 ... [1.0,2.33339,np.nan], decimal=5)
1181 Traceback (most recent call last):
1182 ...
1183 AssertionError:
1184 Arrays are not almost equal to 5 decimals
1185 <BLANKLINE>
1186 Mismatched elements: 1 / 3 (33.3%)
1187 Mismatch at index:
1188 [1]: 2.33333 (ACTUAL), 2.33339 (DESIRED)
1189 Max absolute difference among violations: 6.e-05
1190 Max relative difference among violations: 2.57136612e-05
1191 ACTUAL: array([1. , 2.33333, nan])
1192 DESIRED: array([1. , 2.33339, nan])
1193
1194 >>> np.testing.assert_array_almost_equal([1.0,2.33333,np.nan],
1195 ... [1.0,2.33333, 5], decimal=5)
1196 Traceback (most recent call last):
1197 ...
1198 AssertionError:
1199 Arrays are not almost equal to 5 decimals
1200 <BLANKLINE>
