如何让NumPy不将list的子类视为类数组(array-like)对象?
解决NumPy将list子类视为类数组对象的问题
方案1:预先创建空对象数组再赋值
直接初始化指定形状、dtype=np.object_的空数组,再把子类实例逐个填入。这种方式下NumPy不会尝试解析实例内部结构,只会存储对象引用,完美保留原对象和数组形状。
import numpy as np class A(list): pass values = [A((1,2)), A((3,4))] # 创建对应长度的空object数组 array = np.empty(len(values), dtype=np.object_) # 批量赋值 array[:] = values print(array.dtype) # 输出 object print(type(array[0])) # 输出 <class '__main__.A'>
针对需保持特定形状的场景(比如形状(1,3,1)的Car类数组):
class Car: def __init__(self): self.tensor = np.random.rand(5,3) # 目标形状(1,3,1) target_shape = (1,3,1) car_array = np.empty(target_shape, dtype=np.object_) # 遍历所有索引填充实例 for idx in np.ndindex(target_shape): car_array[idx] = Car() print(car_array.shape) # 输出 (1,3,1) # 单个Car实例的张量形状不受影响 print(car_array[0,0,0].tensor.shape) # 输出 (5,3)
方案2:修改子类,让NumPy不把它当类数组
NumPy判断对象是否为类数组,核心依据是是否实现__array__方法或__array_interface__属性。我们可以给子类重写__array__方法,返回NotImplemented,明确告诉NumPy:这个类不是类数组对象,别展开它。
import numpy as np class A(list): def __array__(self, dtype=None): # 返回NotImplemented,阻止NumPy将该类识别为array-like return NotImplemented values = [A((1,2)), A((3,4))] array = np.array(values, dtype=np.object_) print(array.dtype) # 输出 object print(type(array[0])) # 输出 <class '__main__.A'>
对Car类(哪怕不是list子类,只要不想被NumPy展开)也适用:
class Car: def __init__(self): self.tensor = np.random.rand(5,3) def __array__(self, dtype=None): return NotImplemented # 按目标形状构造嵌套列表 values = [[[Car() for _ in range(1)] for _ in range(3)] for _ in range(1)] car_array = np.array(values, dtype=np.object_) print(car_array.shape) # 输出 (1,3,1)
原代码失效的原因
你写的代码里,虽然指定了dtype=np.object_,但NumPy在处理输入时会先检查元素是否为类数组对象。由于A继承自list,属于NumPy认可的类数组类型,它会自动展开每个A实例内部的元素,合并成一个int类型的数组。只有当元素无法被统一为数值型数组时,NumPy才会回退到存储对象引用的object数组——显然你的例子里两个A实例的内部元素都是int,所以直接被合并了。
内容的提问来源于stack exchange,提问作者Inyoung Kim 김인영
相关产品推荐
相关产品推荐

