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

如何在单元测试中比较浮点数相似性而非严格相等?

解决Flask API单元测试中机器学习浮点数结果的近似断言问题

嘿,这个问题太常见了!机器学习模型因为训练过程中的随机性(比如权重初始化、数据批次的随机划分),每次输出的浮点数预测结果难免有细微差异,直接用self.assertEqual做严格比对肯定会频繁失败。给你两个实用的解决方案,适配你的unittest测试场景:

方案1:逐个字段校验,对目标浮点数用近似断言

既然只有ChangeToSurvive字段有波动,我们可以先校验其他所有字段的严格一致性,再单独对这个浮点数做近似相等断言。unittest自带的assertAlmostEqual方法正好适合这个需求,你可以通过places或delta参数控制误差范围:

修改后的测试代码

def test_by_name(self):
    post_data = {
        'name': 'Andre'
    }
    resp = self.app.post('/survivals', data=json.dumps(post_data), content_type='application/json')
    self.assertEqual(resp.status_code, 200)
    self.assertEqual(resp.content_type, 'application/json')
    content = json.loads(resp.get_data(as_text=True))
    size = len(content['Passengers'])
    self.assertEqual(size, 2)
    self.maxDiff = None
    
    expected_passengers = [
        {
            "SibSp": 1,
            "Sex": "0",
            "PassengerId": 925,
            "Survived": 1,
            "Parch": 2,
            "Age": 1,
            "Name": "Johnston, Mrs. Andrew G (Elizabeth Lily\" Watson)\"",
            "ChangeToSurvive": 74.7,
            "Embarked": "0"
        },
        {
            "SibSp": 0,
            "Sex": "1",
            "PassengerId": 1096,
            "Survived": 0,
            "Parch": 0,
            "Age": 1,
            "Name": "Andrew, Mr. Frank Thomas",
            "ChangeToSurvive": 29.2,
            "Embarked": "0"
        }
    ]
    
    # 遍历每个乘客数据,分开校验
    for actual_passenger, expected_passenger in zip(content['Passengers'], expected_passengers):
        # 复制字典并移除波动字段,校验其他字段严格一致
        actual = actual_passenger.copy()
        expected = expected_passenger.copy()
        actual_change = actual.pop('ChangeToSurvive')
        expected_change = expected.pop('ChangeToSurvive')
        
        self.assertEqual(actual, expected, "非ChangeToSurvive字段数据不匹配")
        
        # 校验浮点数近似相等:根据你的场景选一种方式
        # 方式A:指定小数点后保留n位相等(比如允许0.1的误差)
        self.assertAlmostEqual(actual_change, expected_change, places=1)
        
        # 方式B:指定最大允许差值(比如允许1.5以内的误差,适配你例子中29.2和30.5的情况)
        # self.assertAlmostEqual(actual_change, expected_change, delta=1.5)

关键说明

  • places=1:表示两个浮点数的小数点后第1位必须相等,允许的误差范围是±0.05
  • delta=1.5:表示两个浮点数的差值不超过1.5就算通过,更适合波动较大的场景
  • 先校验其他字段能确保API返回的结构和非预测数据完全正确,避免因其他字段错误导致的测试失败

方案2:固定模型随机种子(可选)

如果你的测试需要完全一致的预测结果,可以在模型训练代码中设置全局随机种子,消除随机性:

import numpy as np
import random
import tensorflow as tf # 如果用TensorFlow/Keras

# 固定所有随机源的种子
random.seed(42)
np.random.seed(42)
tf.random.set_seed(42) # TensorFlow专用

不过这种方式的局限性在于:它会消除模型训练的随机性,可能掩盖某些真实场景下的问题,所以更适合需要稳定测试结果的场合,而近似断言的通用性更强。

内容的提问来源于stack exchange,提问作者Andre Araujo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 07:11:47