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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 01:20:00