【问题标题】:numpy ndarray subclass: ufunc don't return scalar typenumpy ndarray 子类:ufunc 不返回标量类型
【发布时间】:2013-10-13 23:20:30
【问题描述】:

对于numpy.ndarray 子类,ufunc 输出具有相同的类型。这通常很好,但我希望具有标量输出的 ufunc 返回标量类型(例如numpy.float64)。

例子:

import numpy as np

class MyArray(np.ndarray):
    def __new__(cls, array):
        obj = np.asarray(array).view(cls)
        return obj

a = MyArray(np.arange(5))
a*2
# MyArray([0, 2, 4, 6, 8])  => same class as original (i.e. MyArray), ok

a.sum()
# MyArray(10)               => same as original, but here I'd expect np.int64

type(2*a) is type(a.sum())
# True                    
b = a.view(np.ndarray)
type(2*b) is type(b.sum())    
# False

对于标准 numpy 数组,标量输出具有标量类型。那么如何对我的子类有相同的行为呢?

我在 OSX 10.6 上使用 Python 2.7.3 和 numpy 1.6.2

【问题讨论】:

  • 我不确定我是否理解:你想拥有例如sum() 始终返回 float64,无论 a 是整数数组还是浮点数组?
  • 不,我希望它返回与 a.view(np.ndarray).sum() 相同的结果:具有与传统数组相同的行为。
  • 好吧,那很奇怪,因为当我运行你的代码时,a.sum() 对我来说是np.int64。 Python 2.7.5 + NumPy 1.7.0,以及 Python 3.3.2 + Numpy 1.8.0-dev。
  • 注意:您可能会混淆a.sum().dtypetype(a.sum())
  • 我的结果匹配 OP,而不是 @Evert。 Numpy 1.7.1,Python 2.7.5。有趣的是,输出被压缩了,因为它的形状为() 而不是(1,),但它没有转换为标量。因此,如果类型不是ndarraynp.sum() 会忽略keepdims 参数这一事实不会影响。

标签: python numpy subclass scalar multidimensional-array


【解决方案1】:

您需要在 ndarray 子类中使用如下所示的函数覆盖 __array_wrap__

def __array_wrap__(self, obj):
    if obj.shape == ():
        return obj[()]    # if ufunc output is scalar, return it
    else:
        return np.ndarray.__array_wrap__(self, obj)

__array_wrap__ 在 ufunc 之后调用以进行清理工作。在默认实现特殊情况下,精确的 ndarrays(但不是子类)将零秩数组转换为标量。至少对于某些版本的 numpy 来说是这样。

【讨论】:

    猜你喜欢
    • 2016-11-03
    • 2011-09-05
    • 1970-01-01
    • 2022-12-17
    • 1970-01-01
    • 2016-09-13
    • 2019-06-27
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多