使用numpy.testing.assert_array_equal对比含数组字段的结构化数组失败排查
问题:结构化NumPy数组中object dtype元素的断言失败问题
使用numpy.testing.assert_array_equal对比scipy.io.loadmat读取的MAT文件数据与手动生成的预期NumPy数组时,仅变量c的断言失败,变量a、b的断言均正常通过。
MAT文件生成代码
a = [1, 2; 3, 4]; b = struct('MyField', 10); c = struct('MyField', [1, 2; 3, 4]); save('example.mat', 'a', 'b', 'c');
测试代码
import numpy as np from numpy.testing import assert_array_equal from scipy.io import loadmat a = np.array([[1., 2.], [3., 4.]]) b = np.array([[(np.array(10.0),)]], dtype=[("MyField", "O")]) c = np.array( [[ (np.array([[1., 2.], [3., 4.]]),) ]], dtype=[("MyField", "O")]) matdict = loadmat("example.mat", mat_dtype=True) assert_array_equal(matdict["a"], a) # Passes assert_array_equal(matdict["b"], b) # Passes assert_array_equal(matdict["c"], c) # Fails
报错信息
Traceback (most recent call last): File ".../python3.13/site-packages/numpy/testing/_private/utils.py", line 851, in assert_array_compare val = comparison(x, y) ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all() During handling of the above exception, another exception occurred: Traceback (most recent call last): ... File ".../python3.13/site-packages/numpy/testing/_private/utils.py", line 1057, in assert_array_equal assert_array_compare(operator.__eq__, actual, desired, err_msg=err_msg, ~~~~~~~~~~~~~~~~~~~~^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ verbose=verbose, header='Arrays are not equal', ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ strict=strict) ^^^^^^^^^^^^^^ File ".../python3.13/site-packages/numpy/testing/_private/utils.py", line 929, in assert_array_compare raise ValueError(msg) ValueError: error during assertion: Traceback (most recent call last): File ".../python3.13/site-packages/numpy/testing/_private/utils.py", line 851, in assert_array_compare val = comparison(x, y) ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all() Arrays are not equal ACTUAL: array([[(array([[1., 2.], [3., 4.]]),)]], dtype=[('MyField', 'O')]) DESIRED: array([[(array([[1., 2.], [3., 4.]]),)]], dtype=[('MyField', 'O')])
底层原因分析
问题出在assert_array_equal处理object dtype元素的逻辑上:
- 对于变量
b,其object dtype字段里是标量数组np.array(10.0),用==对比时会返回单个布尔值,assert_array_equal可以正常处理。 - 对于变量
c,其object dtype字段里是多维数组,用==对比时会返回一个布尔数组(每个元素对应位置的相等性),而numpy无法直接将布尔数组作为单个布尔值判断,因此抛出"The truth value of an array with more than one element is ambiguous"的错误。
视觉上ACTUAL和DESIRED完全一致,但底层对比逻辑的差异导致了断言失败。
正确的对比方法
方案1:手动遍历元素,逐个对比object内的数组
直接提取结构化数组中object字段的数组,用assert_array_equal逐个验证:
import numpy as np from numpy.testing import assert_array_equal from scipy.io import loadmat # 加载数据 matdict = loadmat("example.mat", mat_dtype=True) # 预期数组 c = np.array( [[ (np.array([[1., 2.], [3., 4.]]),) ]], dtype=[("MyField", "O")]) # 遍历结构化数组的每个元素 actual_c = matdict["c"] for actual_entry, desired_entry in np.nditer([actual_c, desired_c]): # 提取字段内的数组并对比 assert_array_equal(actual_entry["MyField"], desired_entry["MyField"])
方案2:使用numpy.testing.assert_equal替代assert_array_equal
assert_equal会递归处理对象类型的元素,对于object dtype里的多维数组,会自动调用数组的相等判断逻辑,无需手动遍历:
np.testing.assert_equal(matdict["c"], c) # 此断言会正常通过
内容的提问来源于stack exchange,提问作者foreverallama
相关产品推荐
相关产品推荐

