Функция ensure_shape
Функция ensure_shape применяется к тензору и позволяет
утвердить его форму. Первым параметром функция принимает
исходный тензор. Вторым параметром передается ожидаемая
форма в виде кортежа или списка. Если статическая форма
тензора несовместима с указанной, выбрасывается исключение.
Если форма содержит неизвестные размерности, они
уточняются на основе переданного значения.
Синтаксис
tf.ensure_shape(x, shape, [name])
Пример
Давайте создадим тензор и убедимся, что его форма соответствует ожидаемой:
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
@tf.function
def f(x):
return tf.ensure_shape(x, [None, 3])
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = f(t)
print(res.shape)
Результат выполнения кода:
(2, 3)
Пример
Давайте посмотрим, что произойдет при несовпадении формы:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
try:
res = tf.ensure_shape(t, [3])
except ValueError as e:
print("ValueError")
Результат выполнения кода:
"ValueError"
Смотрите также
-
функцию
assert_shapes,
которая проверяет соответствие форм тензоров -
функцию
assert_rank,
которая проверяет ранг тензора -
функцию
tf_function,
которая компилирует функцию в граф TensorFlow -
функцию
assert_same_structure,
которая проверяет одинаковую структуру вложенных объектов