numpy.squeeze()

Функция np.squeeze удаляет из формы массива все оси, длина которых равна 1. Это удобно после операций, добавляющих лишние размерности (например, keepdims=True или загрузка данных с одним каналом).

Сигнатура

numpy.squeeze(a, axis=None)

Параметры

  • a — входной массив.

  • axis — конкретная ось (или кортеж осей) с размером 1, которые нужно удалить. По умолчанию удаляются все.

Возвращаемое значение

Массив с теми же данными, но без осей размерности 1. По возможности — представление (view).

Примеры

Пример 1. Убрать одиночные оси.

import numpy as np

a = np.zeros((1, 3, 1, 4))
b = np.squeeze(a)
print(a.shape, "->", b.shape)
(1, 3, 1, 4) -> (3, 4)

Пример 2. Удалить только одну ось.

a = np.zeros((1, 3, 1, 4))
b = np.squeeze(a, axis=0)
print(b.shape)
(3, 1, 4)

Пример 3. После np.mean с keepdims.

m = np.array([[1, 2, 3], [4, 5, 6]])
means = np.mean(m, axis=1, keepdims=True)
print(means.shape)
print(np.squeeze(means).shape)
(2, 1)
(2,)

Пример 4. Преобразование столбца в одномерный массив.

col = np.array([[10], [20], [30], [40]])
print(col.shape)
print(np.squeeze(col))
(4, 1)
[10 20 30 40]

Пример 5. Ошибка при ``axis`` с размером != 1.

a = np.zeros((1, 3, 1))
try:
    np.squeeze(a, axis=1)
except ValueError as e:
    print("Ошибка:", e)
Ошибка: cannot select an axis to squeeze out which has size not equal to one

См. также

Примечание

Лицензия и источники

Техническое описание функции адаптировано из официальной документации NumPy (https://numpy.org/doc/stable/), BSD-3-Clause License. Примеры и пояснения — © AlashEd Wiki.