【问题标题】:Display the array entry causing a test to fail显示导致测试失败的数组条目
【发布时间】:2020-03-19 09:20:49
【问题描述】:

作为测试套件的一部分,我必须检查函数返回的 numpy 数组是否正确。 使用np.array_equal 很容易进行此检查,它返回一个布尔值,用于判断所有数组元素是否相同。

如果测试失败,则错误消息对于了解导致失败的原因并不是特别有帮助。

import unittest
import numpy as np

class TestArray(unittest.TestCase):
    def test_values(self):
        x = np.array([1, 2])
        self.assertTrue(np.array_equal(x, [1, 3]))


if __name__ == "__main__":
    unittest.main()

测试失败信息:

Traceback (most recent call last):
  File "test.py", line 7, in test_values
    self.assertTrue(np.array_equal(x, [1, 3]))
AssertionError: False is not true

是否有一种简单的方法来检查条目是否相等,即显示第一个不相等条目的索引和值?我想要一条错误消息,例如:

AssertionError: Arrays not equal at index 1 (2 != 3) 

【问题讨论】:

    标签: python python-3.x numpy python-unittest


    【解决方案1】:

    我们可以从np.array_equal获取代码并重写它,在最后添加另一个检查

    def array_equal(a1, a2):
        try:
            a1, a2 = asarray(a1), asarray(a2)
        except Exception:
            return False
        if a1.shape != a2.shape:
            return False
        eq = asarray(a1 == a2) # [ True False False True]
        if not bool(eq.all()):
            errors = [f"idx:{idx} ({vals[0]}!={vals[1]})"
                      for idx, vals in enumerate(zip(a1, a2))
                      if not eq[idx]]
            raise AssertionError("Arrays not equal " + " ".join(errors))
        return True
    
    class TestArray(unittest.TestCase):
        def test_values(self):
            x = np.array([1, 1, 1, 1])
            self.assertTrue(array_equal(x, [1, 2, 3, 1]))
    
    if __name__ == "__main__":
        unittest.main()
    

    AssertionError: Arrays not equal idx:1 (1!=2) idx:2 (1!=3)

    【讨论】:

    • 谢谢,这正是我所追求的。我想知道是否也可以只使用 unittest 和 numpy 函数
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2017-02-18
    • 1970-01-01
    • 2017-07-04
    • 2016-07-21
    相关资源
    最近更新 更多