You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何对包含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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 16:45:06