如何使用mock对比Sample类predict方法返回值且无需加载TensorFlow模型
单元测试修正与mock使用说明
1. 你当前需要加载模型的原因
你误以为实例化Sample类必须传入真实的TensorFlow模型,实际上在这个测试场景下:
- 你要测试的
predict方法仅会调用tf_process并直接返回其结果 - 当你mock了
tf_process方法后,该方法不会执行内部真实逻辑,完全不会用到self.tf_model属性,根本不需要加载真实模型
2. 修正后的测试代码
import unittest import pandas as pd from unittest import mock from my_package.sample_model import Sample class TestSample(unittest.TestCase): def test_predict(self): # 实例化Sample时传入任意占位对象即可,无需真实模型 mock_tf_model = mock.Mock() sample_instance = Sample(tf_model=mock_tf_model) # mock实例的tf_process方法,设置预期返回值 test_input = pd.DataFrame({"feature": [1, 2, 3]}) expected_output = pd.DataFrame({"feature": [1, 2, 3], "tf_predictions": [0.2, 0.7, 0.4]}) sample_instance.tf_process = mock.Mock(return_value=expected_output) # 直接调用predict方法,全程不需要加载真实模型 actual_output = sample_instance.predict(test_input) # 两个校验点覆盖全部逻辑 # 校验tf_process被正确调用且传入了正确参数 sample_instance.tf_process.assert_called_once_with(test_input) # 校验predict返回值和tf_process返回值一致 pd.testing.assert_frame_equal(actual_output, expected_output)
3. 该场景下mock的作用
- 提升测试效率:无需加载大体积模型、无需启动TensorFlow运行时,单条测试执行速度从秒级降到毫秒级
- 保证测试职责单一:该用例仅验证
predict方法的逻辑,tf_process内部的预处理、模型调用逻辑可以单独写测试用例覆盖,避免单条用例耦合过多逻辑 - 提升测试稳定性:不会因为模型文件丢失、TensorFlow版本不兼容等环境问题导致用例失败,只有当
predict方法本身逻辑被修改错误时才会触发报错
内容的提问来源于stack exchange,提问作者data_person
相关产品推荐
相关产品推荐

