numpy的assert_array_equal是否验证复数NaN实虚部匹配的数组相等?
复数NaN数组的相等性验证:
assert_array_equal的局限与自定义实现 首先直接给结论:numpy的assert_array_equal默认无法满足你的需求。原因很简单:numpy中NaN != NaN是成立的,所以默认的断言逻辑会把任何包含NaN的元素判定为不相等——哪怕两个复数的实部和虚部都是完全一样的NaN组合(比如np.nan + 0j和np.nan + 0j)。
要实现你要求的验证逻辑(对应位置要么都是相等的非NaN复数,要么都是NaN复数且实部、虚部分别匹配:实部同为NaN或同为相等数值,虚部同理),你需要自定义检查逻辑。
自定义断言函数实现
这里提供一个针对性的解决方案,核心思路是把复数数组拆解为实部和虚部分别验证,同时处理NaN的特殊情况:
import numpy as np def assert_complex_nan_equal(arr1, arr2): # 第一步:检查数组形状是否一致 np.testing.assert_equal(arr1.shape, arr2.shape, err_msg="Arrays have different shapes") # 拆解实部和虚部 arr1_real = arr1.real arr2_real = arr2.real arr1_imag = arr1.imag arr2_imag = arr2.imag # 验证实部:要么数值相等,要么同为NaN real_valid = np.logical_or( np.equal(arr1_real, arr2_real), np.logical_and(np.isnan(arr1_real), np.isnan(arr2_real)) ) # 验证虚部:同理 imag_valid = np.logical_or( np.equal(arr1_imag, arr2_imag), np.logical_and(np.isnan(arr1_imag), np.isnan(arr2_imag)) ) # 所有位置的实部和虚部都必须满足条件 assert np.all(real_valid), "Mismatch in real parts" assert np.all(imag_valid), "Mismatch in imaginary parts"
测试示例
我们用几种典型情况验证这个函数的行为:
- 正常相等的复数数组:
a = np.array([1+2j, 3+4j]) b = np.array([1+2j, 3+4j]) assert_complex_nan_equal(a, b) # 无报错,验证通过
- 同为NaN复数且实虚部匹配:
a = np.array([np.nan + 0j, 5+np.nan*1j]) b = np.array([np.nan + 0j, 5+np.nan*1j]) assert_complex_nan_equal(a, b) # 无报错,验证通过
- 同为NaN复数但虚部不匹配:
a = np.array([np.nan + 0j]) b = np.array([np.nan + 1j]) assert_complex_nan_equal(a, b) # 抛出AssertionError,符合预期
- 一个是NaN复数,一个是非NaN:
a = np.array([np.nan + 0j]) b = np.array([1+2j]) assert_complex_nan_equal(a, b) # 抛出AssertionError,符合预期
这样就能完全满足你对复数NaN数组的相等性验证要求了。
内容的提问来源于stack exchange,提问作者ARF
相关产品推荐
相关产品推荐

