Написание пользовательских контейнеров массивов
Механизм диспетчеризации NumPy, представленный в версии numpy v1.16, является рекомендуемым подходом для написания пользовательских контейнеров N-мерных массивов, совместимых с API NumPy и предоставляющих пользовательские реализации функциональности NumPy. Примеры включают массивы dask, N-мерный массив, распределенный по нескольким узлам, и массивы cupy, N-мерный массив на графическом процессоре.
Для того чтобы понять, как писать пользовательские контейнеры массивов, начнём с простого примера, который имеет ограниченное применение, но иллюстрирует участвующие в этом концепции.
>>> import numpy as np
>>> class DiagonalArray:
... def __init__(self, N, value):
... self._N = N
... self._i = value
... def __repr__(self):
... return f"{self.__class__.__name__}(N={self._N}, value={self._i})"
... def __array__(self):
... return self._i * np.eye(self._N)
...
Наш пользовательский массив можно инициализировать так:
>>> arr = DiagonalArray(5, 1) >>> arr DiagonalArray(N=5, value=1)
Мы можем преобразовать его в массив NumPy, используя numpy.array или numpy.asarray, что вызовет его метод __array__ для получения стандартного numpy.ndarray.
>>> np.asarray(arr)
array([[1., 0., 0., 0., 0.],
[0., 1., 0., 0., 0.],
[0., 0., 1., 0., 0.],
[0., 0., 0., 1., 0.],
[0., 0., 0., 0., 1.]])
Если мы применим функцию NumPy к arr, NumPy снова воспользуется интерфейсом __array__ для преобразования его в массив и затем применит функцию обычным способом.
>>> np.multiply(arr, 2)
array([[2., 0., 0., 0., 0.],
[0., 2., 0., 0., 0.],
[0., 0., 2., 0., 0.],
[0., 0., 0., 2., 0.],
[0., 0., 0., 0., 2.]])
Обратите внимание, что возвращаемый тип — стандартный numpy.ndarray.
>>> type(arr) numpy.ndarray
Как мы можем передать наш пользовательский тип массива через эту функцию? NumPy позволяет классу указывать, что он хотел бы обрабатывать вычисления пользовательским способом через интерфейсы __array_ufunc__ и __array_function__. Давайте рассмотрим их по отдельности, начиная с _array_ufunc__. Этот метод охватывает Универсальные функции (ufunc), класс функций, который включает, например, numpy.multiply и numpy.sin.
Метод __array_ufunc__ получает:
-
ufunc, функцию, напримерnumpy.multiply -
method, строку, различающуюnumpy.multiply(...)и варианты, такие какnumpy.multiply.outer,numpy.multiply.accumulate, и так далее. В общем случае этоnumpy.multiply(...),method == '__call__'. -
inputs, который может быть смесью различных типов -
kwargs, ключевые аргументы, передаваемые функции
В этом примере мы будем обрабатывать только метод __call__.
>>> from numbers import Number
>>> class DiagonalArray:
... def __init__(self, N, value):
... self._N = N
... self._i = value
... def __repr__(self):
... return f"{self.__class__.__name__}(N={self._N}, value={self._i})"
... def __array__(self):
... return self._i * np.eye(self._N)
... def __array_ufunc__(self, ufunc, method, *inputs, **kwargs):
... if method == '__call__':
... N = None
... scalars = []
... for input in inputs:
... if isinstance(input, Number):
... scalars.append(input)
... elif isinstance(input, self.__class__):
... scalars.append(input._i)
... if N is not None:
... if N != self._N:
... raise TypeError("inconsistent sizes")
... else:
... N = self._N
... else:
... return NotImplemented
... return self.__class__(N, ufunc(*scalars, **kwargs))
... else:
... return NotImplemented
...
Теперь наш пользовательский тип массива проходит через функции NumPy.
>>> arr = DiagonalArray(5, 1) >>> np.multiply(arr, 3) DiagonalArray(N=5, value=3) >>> np.add(arr, 3) DiagonalArray(N=5, value=4) >>> np.sin(arr) DiagonalArray(N=5, value=0.8414709848078965)
На данном этапе arr + 3 не работает.
>>> arr + 3 TypeError: unsupported operand type(s) for *: 'DiagonalArray' and 'int'
Для его поддержки нам необходимо определить Python-интерфейсы __add__, __lt__, и так далее, чтобы выполнить диспетчеризацию к соответствующей универсальной функции. Это удобно сделать, унаследовав от миксина NDArrayOperatorsMixin.
>>> import numpy.lib.mixins
>>> class DiagonalArray(numpy.lib.mixins.NDArrayOperatorsMixin):
... def __init__(self, N, value):
... self._N = N
... self._i = value
... def __repr__(self):
... return f"{self.__class__.__name__}(N={self._N}, value={self._i})"
... def __array__(self):
... return self._i * np.eye(self._N)
... def __array_ufunc__(self, ufunc, method, *inputs, **kwargs):
... if method == '__call__':
... N = None
... scalars = []
... for input in inputs:
... if isinstance(input, Number):
... scalars.append(input)
... elif isinstance(input, self.__class__):
... scalars.append(input._i)
... if N is not None:
... if N != self._N:
... raise TypeError("inconsistent sizes")
... else:
... N = self._N
... else:
... return NotImplemented
... return self.__class__(N, ufunc(*scalars, **kwargs))
... else:
... return NotImplemented
...
>>> arr = DiagonalArray(5, 1) >>> arr + 3 DiagonalArray(N=5, value=4) >>> arr > 0 DiagonalArray(N=5, value=True)
Теперь давайте займёмся __array_function__. Мы создадим словарь, сопоставляющий функции NumPy с нашими пользовательскими вариантами.
>>> HANDLED_FUNCTIONS = {}
>>> class DiagonalArray(numpy.lib.mixins.NDArrayOperatorsMixin):
... def __init__(self, N, value):
... self._N = N
... self._i = value
... def __repr__(self):
... return f"{self.__class__.__name__}(N={self._N}, value={self._i})"
... def __array__(self):
... return self._i * np.eye(self._N)
... def __array_ufunc__(self, ufunc, method, *inputs, **kwargs):
... if method == '__call__':
... N = None
... scalars = []
... for input in inputs:
... # In this case we accept only scalar numbers or DiagonalArrays.
... if isinstance(input, Number):
... scalars.append(input)
... elif isinstance(input, self.__class__):
... scalars.append(input._i)
... if N is not None:
... if N != self._N:
... raise TypeError("inconsistent sizes")
... else:
... N = self._N
... else:
... return NotImplemented
... return self.__class__(N, ufunc(*scalars, **kwargs))
... else:
... return NotImplemented
... def __array_function__(self, func, types, args, kwargs):
... if func not in HANDLED_FUNCTIONS:
... return NotImplemented
... # Note: this allows subclasses that don't override
... # __array_function__ to handle DiagonalArray objects.
... if not all(issubclass(t, self.__class__) for t in types):
... return NotImplemented
... return HANDLED_FUNCTIONS[func](*args, **kwargs)
...
Удобный шаблон — определение декоратора implements, который можно использовать для добавления функций к HANDLED_FUNCTIONS.
>>> def implements(np_function): ... "Register an __array_function__ implementation for DiagonalArray objects." ... def decorator(func): ... HANDLED_FUNCTIONS[np_function] = func ... return func ... return decorator ...
Теперь мы напишем реализации функций NumPy для DiagonalArray. Для полноты, чтобы обеспечить использование arr.sum(), добавим метод sum, который вызывает numpy.sum(self), и аналогично для mean.
>>> @implements(np.sum) ... def sum(arr): ... "Implementation of np.sum for DiagonalArray objects" ... return arr._i * arr._N ... >>> @implements(np.mean) ... def mean(arr): ... "Implementation of np.mean for DiagonalArray objects" ... return arr._i / arr._N ... >>> arr = DiagonalArray(5, 1) >>> np.sum(arr) 5 >>> np.mean(arr) 0.2
Если пользователь попытается использовать какие-либо функции NumPy, не включённые в HANDLED_FUNCTIONS, NumPy поднимет исключение TypeError, указывая, что эта операция не поддерживается. Например, конкатенация двух DiagonalArrays не даёт другой диагональный массив, поэтому она не поддерживается.
>>> np.concatenate([arr, arr]) TypeError: no implementation found for 'numpy.concatenate' on types that implement __array_function__: [<class '__main__.DiagonalArray'>]
Кроме того, наши реализации sum и mean не принимают необязательные аргументы, которые есть у реализации NumPy.
>>> np.sum(arr, axis=0) TypeError: sum() got an unexpected keyword argument 'axis'
Пользователь всегда может преобразовать в обычный numpy.ndarray с помощью numpy.asarray и использовать стандартный NumPy оттуда.
>>> np.concatenate([np.asarray(arr), np.asarray(arr)])
array([[1., 0., 0., 0., 0.],
[0., 1., 0., 0., 0.],
[0., 0., 1., 0., 0.],
[0., 0., 0., 1., 0.],
[0., 0., 0., 0., 1.],
[1., 0., 0., 0., 0.],
[0., 1., 0., 0., 0.],
[0., 0., 1., 0., 0.],
[0., 0., 0., 1., 0.],
[0., 0., 0., 0., 1.]])
Обратитесь к исходному коду dask и исходному коду cupy для более подробных примеров пользовательских контейнеров массивов.
См. также NEP 18.
© 2005–2020 NumPy Developers
Licensed under the 3-clause BSD License.
https://numpy.org/doc/1.19/user/basics.dispatch.html