Класс Function
Класс Function в TensorFlow представляет собой обертку для Python-функций,
которая позволяет компилировать их в графовые вычисления.
Первый параметр принимает Python-функцию для обертки.
Второй параметр input_signature задает сигнатуру входных тензоров.
Третий параметр autograph управляет автоматическим преобразованием кода.
Четвертый параметр reduce_retracing управляет повторной трассировкой.
Синтаксис
tf.function(func, [input_signature], [autograph], [reduce_retracing])
Пример
Давайте создадим простую функцию и обернем ее с помощью tf.function:
import tensorflow as tf
@tf.function
def add_numbers(a, b):
return a + b
t1 = tf.constant(2)
t2 = tf.constant(3)
res = add_numbers(t1, t2)
print(res)
Результат выполнения кода:
tf.Tensor(5, shape=(), dtype=int32)
Пример
Давайте создадим Function с явной сигнатурой входных тензоров:
import tensorflow as tf
def multiply_by_two(x):
return x * 2
func = tf.function(
multiply_by_two,
input_signature=[tf.TensorSpec(shape=[None], dtype=tf.int32)]
)
t = tf.constant([1, 2, 3, 4, 5])
res = func(t)
print(res)
Результат выполнения кода:
tf.Tensor([ 2 4 6 8 10], shape=(5,), dtype=int32)
Пример
Давайте получим конкретную функцию с помощью метода get_concrete_function:
import tensorflow as tf
@tf.function
def square(x):
return x * x
concrete = square.get_concrete_function(
tf.TensorSpec(shape=[None], dtype=tf.int32)
)
print(concrete)
Результат выполнения кода:
ConcreteFunction square(x)
Args:
x: TensorSpec(shape=(None,), dtype=tf.int32, name='x')
Returns:
TensorSpec(shape=(None,), dtype=tf.int32, name=None)
Пример
Давайте посмотрим на атрибуты класса Function:
import tensorflow as tf
@tf.function
def process(x):
return x + 1
res = process(tf.constant(1))
print(process.input_signature)
print(type(process))
Результат выполнения кода:
None
<class 'tensorflow.python.eager.polymorphic_function.polymorphic_function.Function'>
Смотрите также
-
класс
Function,
который представляет обертку для функций -
метод
get_concrete_function,
который возвращает конкретную функцию -
метод
experimental_get_compiler_ir,
который возвращает промежуточное представление компилятора