tf.ensure_shape
| Просмотреть исходный код на GitHub |
Обновляет форму тензора и проверяет во время выполнения, что форма сохраняется.
tf.ensure_shape(
x, shape, name=None
)
При выполнении эта операция утверждает, что форма входного тензора x совместима с аргументом shape. См. tf.TensorShape.is_compatible_with для получения подробной информации.
x = tf.constant([[1, 2, 3],
[4, 5, 6]])
x = tf.ensure_shape(x, [2, 3])
Используйте None для неизвестных измерений:
x = tf.ensure_shape(x, [None, 3]) x = tf.ensure_shape(x, [2, None])
Если форма тензора не совместима с аргументом shape, возникает ошибка:
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), tf.ensure_shape обновляет статическую форму тензора результата, объединяя две формы. См. tf.TensorShape.merge_with для получения подробной информации.
Это наиболее полезно, когда вам известна форма, которая не может быть статически определена TensorFlow.
Следующая тривиальная tf.function выводит статическую форму входного тензора до и после применения ensure_shape.
@tf.function
def f(tensor):
print("Static-shape before:", tensor.shape)
tensor = tf.ensure_shape(tensor, [None, 3])
print("Static-shape after:", tensor.shape)
return tensor
Это позволяет увидеть эффект tf.ensure_shape при трассировке функции:
>>> cf = f.get_concrete_function(tf.TensorSpec([None, None])) Static-shape before: (None, None) Static-shape after: (None, 3)
cf(tf.zeros([3, 3])) # Passes cf(tf.constant([1, 2, 3])) # fails Traceback (most recent call last): InvalidArgumentError: Shape of tensor x [3] is not compatible with expected shape [3,3].
В приведенном выше примере возникает tf.errors.InvalidArgumentError, потому что форма x, (3,), не совместима с аргументом shape, (None, 3)
В контексте 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. |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/ensure_shape