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.