如何对包含NaN的列表编写正确的相等性测试?
解决含NaN序列的unittest断言失败问题
问题说明
自定义holder类的calculate_rolling_average方法计算滚动平均值后,返回的列表包含NaN值。由于NaN与任何值(包括自身)都不相等,使用unittest的assertSequenceEqual断言时,即使两个序列的元素看起来完全一致,测试也会失败。
代码示例
from math import nan import unittest import pandas as pd class holder: def __init__(self, source: list) -> None: self.source = source self.data = [] def calculate_rolling_average(self): self.data = pd.Series(self.source).rolling(3).mean().to_list() class test_value(unittest.TestCase): def setUp(self): self.hold = holder(source=[10, 20, 30, 40]) def test_rolling_simple_average_for_list(self): expected_result = [nan, nan, 20.0, 30.0] self.hold.calculate_rolling_average() self.assertSequenceEqual(self.hold.data, expected_result) unittest.main()
报错信息
AssertionError: Sequences differ: [nan, nan, 20.0, 30.0] != [nan, nan, 20.0, 30.0] First differing element 0: nan nan [nan, nan, 20.0, 30.0]
解决方案
方案1:利用pandas的断言工具(推荐,适配项目场景)
因为代码已经使用pandas计算滚动平均值,直接用pandas提供的assert_series_equal方法,它会自动处理NaN的相等性判断:
import pandas.testing as tm class test_value(unittest.TestCase): def setUp(self): self.hold = holder(source=[10, 20, 30, 40]) def test_rolling_simple_average_for_list(self): expected_result = [nan, nan, 20.0, 30.0] self.hold.calculate_rolling_average() tm.assert_series_equal(pd.Series(self.hold.data), pd.Series(expected_result))
方案2:自定义断言函数(无额外依赖)
如果不想引入pandas测试工具,可自己实现一个能处理NaN的序列断言方法:
class test_value(unittest.TestCase): def setUp(self): self.hold = holder(source=[10, 20, 30, 40]) def assert_sequence_with_nan_equal(self, seq1, seq2): # 先检查序列长度一致 self.assertEqual(len(seq1), len(seq2)) # 遍历每个元素逐一判断 for a, b in zip(seq1, seq2): # 若两边都是NaN则跳过,否则断言值相等 if pd.isna(a) and pd.isna(b): continue self.assertEqual(a, b) def test_rolling_simple_average_for_list(self): expected_result = [nan, nan, 20.0, 30.0] self.hold.calculate_rolling_average() self.assert_sequence_with_nan_equal(self.hold.data, expected_result)
方案3:使用numpy的数组断言
如果项目使用numpy,可将列表转为numpy数组,利用array_equal的equal_nan参数:
import numpy as np class test_value(unittest.TestCase): def setUp(self): self.hold = holder(source=[10, 20, 30, 40]) def test_rolling_simple_average_for_list(self): expected_result = [nan, nan, 20.0, 30.0] self.hold.calculate_rolling_average() self.assertTrue(np.array_equal(self.hold.data, expected_result, equal_nan=True))
内容的提问来源于stack exchange,提问作者Charizard_knows_to_code
相关产品推荐
相关产品推荐

