【问题标题】:Subclass of numpy ndarray doesn't work as expectednumpy ndarray 的子类不能按预期工作
【发布时间】:2014-05-19 10:09:43
【问题描述】:

`大家好。

我发现子类化 ndarray 时有一个奇怪的行为。

import numpy as np

class fooarray(np.ndarray):
    def __new__(cls, input_array, *args, **kwargs):
        obj = np.asarray(input_array).view(cls)
        return obj

    def __init__(self, *args, **kwargs):
        return

    def __array_finalize__(self, obj):
        return

a=fooarray(np.random.randn(3,5))
b=np.random.randn(3,5)

a_sum=np.sum(a,axis=0,keepdims=True)
b_sum=np.sum(b,axis=0, keepdims=True)

print a_sum.ndim #1
print b_sum.ndim #2

如您所见,keepdims 参数不适用于我的子类fooarray。它失去了一根轴。我怎么不能避免这个问题?或者更一般地说,我怎样才能正确地继承 numpy ndarray?

【问题讨论】:

  • 一种可能的解决方案是改用a.sum。

标签: python numpy


【解决方案1】:

np.sum 可以接受各种对象作为输入:不仅是 ndarray,还包括列表、生成器、np.matrixs,例如。 keepdims 参数显然对列表或生成器没有意义。它也不适合np.matrix 实例,因为np.matrixs 总是有2 个维度。如果您查看np.matrix.sum 的调用签名,您会发现它的sum 方法没有keepdims 参数:

Definition: np.matrix.sum(self, axis=None, dtype=None, out=None)

所以ndarray 的一些子类可能有sum 方法,而这些方法没有keepdims 参数。这是对Liskov substitution principle 的不幸违反以及您遇到的陷阱的根源。

现在,如果您查看the source code for np.sum,您会发现它是一个委托函数,它试图根据第一个参数的类型来确定要做什么。

如果第一个参数的类型不是ndarray,则删除keepdims 参数。这样做是因为将 keepdims 参数传递给 np.matrix.sum 会引发异常。

因此,因为np.sum 试图以最一般的方式进行委托,而不是对 ndarray 的子类可能采用的参数做任何假设,所以它在传递 fooarray 时删除了 keepdims 参数。

解决方法是不使用np.sum,而是调用a.sum。无论如何,这更直接,因为np.sum 只是一个委托函数。

import numpy as np


class fooarray(np.ndarray):
    def __new__(cls, input_array, *args, **kwargs):
        obj = np.asarray(input_array, *args, **kwargs).view(cls)
        return obj

a = fooarray(np.random.randn(3, 5))
b = np.random.randn(3, 5)

a_sum = a.sum(axis=0, keepdims=True)
b_sum = np.sum(b, axis=0, keepdims=True)

print(a_sum.ndim)  # 2
print(b_sum.ndim)  # 2

【讨论】:

  • 完全回答我的问题。谢谢=]
【解决方案2】:

详细说明@mskimm 的评论,如果你看看相关的 numpy 源代码的一部分,core/fromnumeric.py,原因很清楚 a.sum(..., keepdims=True) 有效,而 np.sum(a, ..., keepdims=True) 不会:

def sum(a, axis=None, dtype=None, out=None, keepdims=False):
    ...
    if isinstance(a, _gentype):
        res = _sum_(a)
        if out is not None:
            out[...] = res
            return out
        return res
    elif type(a) is not mu.ndarray:
        try:
            sum = a.sum
        except AttributeError:
            return _methods._sum(a, axis=axis, dtype=dtype,
                                out=out, keepdims=keepdims)
        # NOTE: Dropping the keepdims parameters here...
        return sum(axis=axis, dtype=dtype, out=out)
    else:
        return _methods._sum(a, axis=axis, dtype=dtype,
                            out=out, keepdims=keepdims)
    ...

由于您已将np.ndarray 子类化,因此type(a) 是fooarray,而不是 mu.ndarray,所以你最终在这一行:

# NOTE: Dropping the keepdims parameters here...
return sum(axis=axis, dtype=dtype, out=out)

keepdims 关键字参数是ndarrays 的一个相对较新的功能,目前还没有为某些其他类似数组的类实现,例如np.matrix 或np.ma.masked_array,它们也有一个.sum() 方法,因此为什么该参数当前被非ndarrays 删除。

【讨论】:

    猜你喜欢
    • 2022-10-12
    • 1970-01-01
    • 2015-06-26
    • 2013-05-28
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多