Функция squeeze
Функция squeeze удаляет из тензора все оси,
размер которых равен 1. Первым параметром
функция принимает исходный тензор. Вторым параметром
можно передать конкретную ось axis, которую
нужно удалить. Если ось не указана, удаляются все
оси размером 1. Если указанная ось имеет
размер больше 1, возникнет ошибка.
Синтаксис
tf.squeeze(input, [axis])
Пример
Давайте создадим тензор формы (1, 5, 1) и
удалим все оси размером 1:
import tensorflow as tf
t = tf.constant([[[1], [2], [3], [4], [5]]])
res = tf.squeeze(t)
print(t.shape)
print(res.shape)
print(res)
Результат выполнения кода:
(1, 5, 1)
(5,)
tf.Tensor([1 2 3 4 5], shape=(5,), dtype=int32)
Пример
Давайте удалим только конкретную ось, передав ее
номер в параметр axis:
import tensorflow as tf
t = tf.constant([[[1], [2], [3], [4], [5]]])
res = tf.squeeze(t, axis=0)
print(res.shape)
print(res)
Результат выполнения кода:
(5, 1)
tf.Tensor(
[[1]
[2]
[3]
[4]
[5]], shape=(5, 1), dtype=int32)
Пример
Давайте попробуем удалить ось, размер которой
не равен 1, и посмотрим на ошибку:
import tensorflow as tf
t = tf.constant([[[1], [2], [3], [4], [5]]])
res = tf.squeeze(t, axis=1)
print(res.shape)
Результат выполнения кода:
"ValueError: Cannot squeeze axis 1 with size 5"
Смотрите также
-
функцию
expand_dims,
которая добавляет к тензору новую ось -
функцию
reshape,
которая изменяет форму тензора -
функцию
shape,
которая возвращает форму тензора -
функцию
rank,
которая возвращает размерность тензора