tf.ensure_shape
| Просмотреть исходный код на GitHub |
Обновляет форму тензора и проверяет во время выполнения, что форма соответствует.
tf.ensure_shape(
x, shape, name=None
)
При выполнении eager это утверждение формы, которое возвращает входные данные:
x = tf.constant([1,2,3]) print(x.shape) (3,) x = tf.ensure_shape(x, [3]) x = tf.ensure_shape(x, [5]) Traceback (most recent call last): tf.errors.InvalidArgumentError: Shape of tensor dummy_input [3] is not compatible with expected shape [5]. [Op:EnsureShape]
Внутри контекста tf.function или v1.Graph он проверяет форму во время построения и во время выполнения. Это более строго, чем tf.Tensor.set_shape, которое проверяет только форму во время построения.
Примечание: Это отличается отtf.Tensor.set_shapeтем, что устанавливает статическую форму результирующего тензора и применяет её во время выполнения, вызывая ошибку, если форма тензора во время выполнения несовместима со заданной формой.tf.Tensor.set_shapeустанавливает статическую форму тензора без применения её во время выполнения, что может привести к несоответствиям между статически известной формой тензоров и значением тензоров во время выполнения.
Например, при загрузке изображений известного размера:
@tf.function
def decode_image(png):
image = tf.image.decode_png(png, channels=3)
# the `print` executes during tracing.
print("Initial shape: ", image.shape)
image = tf.ensure_shape(image,[28, 28, 3])
print("Final shape: ", image.shape)
return image
При прослеживании функции операции не выполняются, формы могут быть неизвестны. Подробности см. в Руководстве по конкретным функциям.
concrete_decode = decode_image.get_concrete_function(
tf.TensorSpec([], dtype=tf.string))
Initial shape: (None, None, 3)
Final shape: (28, 28, 3)
image = tf.random.uniform(maxval=255, shape=[28, 28, 3], dtype=tf.int32) image = tf.cast(image,tf.uint8) png = tf.image.encode_png(image) image2 = concrete_decode(png) print(image2.shape) (28, 28, 3)
image = tf.concat([image,image], axis=0) print(image.shape) (56, 28, 3) png = tf.image.encode_png(image) image2 = concrete_decode(png) Traceback (most recent call last): tf.errors.InvalidArgumentError: Shape of tensor DecodePng [56,28,3] is not compatible with expected shape [28,28,3].
@tf.function
def bad_decode_image(png):
image = tf.image.decode_png(png, channels=3)
# the `print` executes during tracing.
print("Initial shape: ", image.shape)
# BAD: forgot to use the returned tensor.
tf.ensure_shape(image,[28, 28, 3])
print("Final shape: ", image.shape)
return image
image = bad_decode_image(png) Initial shape: (None, None, 3) Final shape: (None, None, 3) print(image.shape) (56, 28, 3)
| Аргументы | |
|---|---|
x | A Tensor. |
shape | A TensorShape представляющая форму этого тензора, a TensorShapeProto, список, кортеж или None. |
name | Имя этой операции (необязательно). По умолчанию "EnsureShape". |
| Возвращает | |
|---|---|
A Tensor. Имеет тот же тип и содержимое, что и x. Во время выполнения, вызывает tf.errors.InvalidArgumentError, если shape несовместима с формой x. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.3/api_docs/python/tf/ensure_shape