Функция ensure_shape
Функция ensure_shape применяется к тензору и гарантирует, что его форма
соответствует ожидаемой. Первым параметром функция принимает тензор, форму
которого нужно проверить. Вторым параметром передается ожидаемая форма
в виде кортежа или списка. Если форма тензора несовместима с указанной,
функция возбуждает исключение. Если форма тензора известна лишь частично,
функция устанавливает недостающие размерности.
Функция полезна при построении моделей, когда нужно явно указать ожидаемую форму данных, а также для отладки кода.
Синтаксис
tf.ensure_shape(input, shape)
Пример
Давайте создадим тензор и проверим, что его форма равна 5:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.ensure_shape(t, (5,))
print(res)
Результат выполнения кода:
tf.Tensor([1 2 3 4 5], shape=(5,), dtype=int32)
Пример
Давайте попробуем передать несовместимую форму и посмотрим, что произойдет:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.ensure_shape(t, (3,))
print(res)
Результат выполнения кода:
ValueError: Input tensor shape (5,) is not compatible with expected shape (3,)
Пример
Давайте установим форму для тензора, у которого одна из размерностей неизвестна:
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.ensure_shape(t, (2, None))
print(res)
Результат выполнения кода:
tf.Tensor(
[[1 2 3]
[4 5 6]], shape=(2, 3), dtype=int32)
Пример
Давайте проверим форму тензора, созданного функцией zeros:
import tensorflow as tf
t = tf.zeros((3, 4))
res = tf.ensure_shape(t, (3, 4))
print(res)
Результат выполнения кода:
tf.Tensor(
[[0. 0. 0. 0.]
[0. 0. 0. 0.]
[0. 0. 0. 0.]], shape=(3, 4), dtype=float32)
Смотрите также
-
функцию
shape,
которая возвращает форму тензора -
функцию
reshape,
которая изменяет форму тензора -
функцию
expand_dims,
которая добавляет новую размерность -
функцию
squeeze,
которая удаляет размерности единичной длины