NumPy: вычисления и устройство массивов

Агрегации по массиву и по оси: sum, mean, min, max, argmin, argmax

Содержание курса

Параметр axis: какая ось «схлопывается» и как меняется shape

Когда нужна не одна общая цифра, а результат отдельно по каждой строке или по каждому столбцу, в игру вступает параметр axis. Его смысл проще всего описать через одно слово: схлопывание. Указывая axis=k, вы говорите NumPy: «пробегись вдоль оси с номером k и сверни её». Ось исчезает из результата — в этом и есть весь механизм.

Разберём на двумерном массиве shape (3, 3). Оси здесь две: ось 0 — это направление «вниз по строкам», ось 1 — «вправо по столбцам».

import numpy as np

arr = np.array([[3, 1, 4],
                [1, 5, 9],
                [2, 6, 5]])
# shape: (3, 3)

print(np.sum(arr, axis=0))  # [6, 12, 18] — shape (3,)
print(np.sum(arr, axis=1))  # [8, 15, 13] — shape (3,)

С axis=0 NumPy суммирует вдоль строк: берёт первый столбец [3, 1, 2], складывает в 6; второй [1, 5, 6]12; третий [4, 9, 5]18. Ось 0 (размерность «три строки») исчезла — остался массив из трёх чисел, по одному на столбец. Shape исходника был (3, 3), вычеркиваем нулевой элемент — получаем (3,).

С axis=1 всё симметрично: суммируем вдоль столбцов, то есть каждую строку сворачиваем в одно число. Ось 1 исчезает, результат тоже (3,) — по одному числу на строку.

Общее правило предсказания shape:

  • Запишите shape исходного массива как кортеж.
  • Уберите из кортежа элемент с позицией k.
  • То, что осталось, и есть shape результата.

Это правило работает для массивов любой размерности. Например, массив shape (2, 3, 4) при axis=1 даёт (2, 4): элемент с позицией 1 (то есть 3) вычеркнут.

arr3d = np.zeros((2, 3, 4))
result = np.mean(arr3d, axis=1)
print(result.shape)  # (2, 4)

Практически это означает: если хотите получить среднее по каждому столбцу таблицы — берёте axis=0; среднее по каждой строке — axis=1. Правило «убираем k-й элемент из кортежа» позволяет не держать в голове конкретные числа, а вычислять результат из первых принципов для любого массива.

Обратите внимание: axis работает одинаково для всех агрегирующих функций — np.sum, np.mean, np.min, np.max и, как увидим дальше, np.argmin, np.argmax. Параметр описывает только то, какая ось схлопывается, а не какая операция выполняется.