如何为perfplot自定义equality_check相等检查函数?
解决perfplot自定义相等检查的问题
你需要编写一个自定义相等检查函数,适配返回包含不同长度numpy数组的列表这类场景。这个函数需要完成两个核心步骤:
- 先验证两个结果的列表长度是否一致
- 逐个遍历列表元素,对每个numpy数组执行相等判断(根据需求选择严格相等或允许浮点误差的检查)
以下是修改后的完整可运行代码:
import perfplot import numpy as np def f1(rng): return [np.array(range(i)) for i in rng] def f2(rng): return [np.array(list(rng)[:i]) for i in range(len(rng))] def f3(rng): return [np.array([*range(i)]) for i in rng] # 自定义相等检查函数 def custom_equality_check(a, b): # 先检查列表长度是否匹配 if len(a) != len(b): return False # 逐个对比列表中的numpy数组 for arr_a, arr_b in zip(a, b): # 用np.array_equal做严格相等校验;若需允许浮点误差,替换为np.allclose即可 if not np.array_equal(arr_a, arr_b): return False return True perfplot.show( kernels=[f1, f2, f3], n_range=[10**k for k in range(4)], setup=lambda n: range(n), equality_check=custom_equality_check # 传入自定义检查函数 )
补充说明
- 若你的函数返回浮点型数组,需要允许计算误差的话,直接把
np.array_equal替换为np.allclose即可 - 自定义检查函数必须接收两个参数(待对比的两个结果),最终返回布尔值表示是否相等
内容的提问来源于stack exchange,提问作者cottontail
相关产品推荐
相关产品推荐

