Функция debugging.assert_shapes
Функция debugging.assert_shapes применяется к тензорам
для проверки их форм во время выполнения программы.
Первым параметром функция принимает список пар,
где каждая пара состоит из тензора и ожидаемой формы.
Вторым параметром можно передать сообщение об ошибке.
Если форма тензора не совпадает с ожидаемой,
функция выбрасывает исключение.
Синтаксис
tf.debugging.assert_shapes(shapes, [message])
Пример
Давайте проверим форму тензора t,
которая должна быть равна (5,):
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.debugging.assert_shapes([(t, (5,))])
print(res)
Результат выполнения кода:
None
Пример
Давайте проверим форму двумерного тензора t,
которая должна быть равна (2, 3):
import tensorflow as tf
t = tf.constant([[1, 2, 3], [4, 5, 6]])
res = tf.debugging.assert_shapes([(t, (2, 3))])
print(res)
Результат выполнения кода:
None
Пример
Давайте передадим неверную ожидаемую форму и посмотрим, какая ошибка возникнет:
import tensorflow as tf
t = tf.constant([1, 2, 3, 4, 5])
res = tf.debugging.assert_shapes([(t, (4,))])
print(res)
Результат выполнения кода:
"InvalidArgumentError: Dimensions must be equal, but are 5 and 4"
Смотрите также
-
функцию
debugging.assert_equal,
которая проверяет равенство двух тензоров -
функцию
debugging.assert_rank,
которая проверяет ранг тензора -
функцию
debugging.assert_type,
которая проверяет тип данных тензора -
функцию
ensure_shape,
которая устанавливает форму тензора