Метод get_concrete_function
Метод get_concrete_function принадлежит классу
Function и позволяет получить конкретную
функцию - объект ConcreteFunction, в котором
зафиксированы типы и формы входных тензоров.
Метод вызывается у объекта tf.function и
принимает на вход те же аргументы, что и сама
функция. Для каждого набора аргументов TensorFlow
строит отдельный граф, а метод возвращает ссылку
на этот граф в виде конкретной функции.
Первым параметром метод принимает аргументы
обернутой функции - числа, тензоры или объекты
tf.TensorSpec. Дополнительно можно передать
именованные аргументы, соответствующие параметрам
исходной функции.
Синтаксис
f.get_concrete_function(*args, **kwargs)
Пример
Давайте создадим объект tf.function и получим
для него конкретную функцию с тензором
1, 2, 3, 4, 5:
import tensorflow as tf
@tf.function
def add_one(x):
return x + 1
concrete = add_one.get_concrete_function(
tf.constant([1, 2, 3, 4, 5])
)
print(type(concrete))
Результат выполнения кода:
<class 'tensorflow.python.eager.polymorphic_function.polymorphic_function.ConcreteFunction'>
Пример
Давайте вызовем конкретную функцию и получим результат ее работы:
import tensorflow as tf
@tf.function
def add_one(x):
return x + 1
concrete = add_one.get_concrete_function(
tf.constant([1, 2, 3, 4, 5])
)
res = concrete(tf.constant([1, 2, 3, 4, 5]))
print(res)
Результат выполнения кода:
tf.Tensor([2 3 4 5 6], shape=(5,), dtype=int32)
Пример
Давайте передадим в метод get_concrete_function
описание входного тензора через tf.TensorSpec:
import tensorflow as tf
@tf.function
def add_one(x):
return x + 1
spec = tf.TensorSpec(shape=(5,), dtype=tf.int32)
concrete = add_one.get_concrete_function(spec)
res = concrete(tf.constant([1, 2, 3, 4, 5]))
print(res)
Результат выполнения кода:
tf.Tensor([2 3 4 5 6], shape=(5,), dtype=int32)
Пример
Давайте посмотрим на структуру графа конкретной
функции через атрибут graph:
import tensorflow as tf
@tf.function
def add_one(x):
return x + 1
concrete = add_one.get_concrete_function(
tf.constant([1, 2, 3, 4, 5])
)
print(concrete.graph.get_operations())
Результат выполнения кода:
[<tf.Operation 'x' type=Placeholder>, <tf.Operation 'add/y' type=Const>, <tf.Operation 'add' type=AddV2>, <tf.Operation 'Identity' type=Identity>]
Смотрите также
-
класс
Function,
который оборачивает Python-функцию в граф TensorFlow -
метод
get_concrete_function,
который возвращает конкретную функцию для заданных входов -
метод
experimental_get_compiler_ir,
который возвращает промежуточное представление скомпилированной функции