Функция tf.function
Функция tf.function преобразует обычную
Python-функцию в графовую TensorFlow-функцию.
Первым параметром передаётся функция, которую
нужно скомпилировать. Вторым параметром можно
передать список входных сигнатур input_signature.
Также доступны параметры reduce_retracing и
autograph. Скомпилированная функция выполняется
быстрее за счёт построения статического графа и
может быть экспортирована в SavedModel.
Синтаксис
tf.function(func, [input_signature], [reduce_retracing], [autograph])
Пример
Давайте скомпилируем простую функцию сложения двух тензоров с помощью декоратора:
import tensorflow as tf
@tf.function
def add(a, b):
return a + b
res = add(tf.constant([1, 2, 3, 4, 5]), tf.constant([1, 2, 3, 4, 5]))
print(res)
Результат выполнения кода:
tf.Tensor([ 2 4 6 8 10], shape=(5,), dtype=int32)
Пример
Давайте применим tf.function без декоратора
и передадим функцию как аргумент:
import tensorflow as tf
def multiply(a, b):
return a * b
compiled = tf.function(multiply)
res = compiled(tf.constant([[1, 2, 3], [4, 5, 6]]), tf.constant(2))
print(res)
Результат выполнения кода:
tf.Tensor(
[[ 2 4 6]
[ 8 10 12]], shape=(2, 3), dtype=int32)
Пример
Давайте укажем входную сигнатуру через параметр
input_signature:
Результат выполнения кода:
tf.Tensor([ 2 4 6 8 10], shape=(5,), dtype=int32)
Пример
Давайте посмотрим на граф скомпилированной функции
через атрибут graph:
import tensorflow as tf
@tf.function
def square(t):
return t * t
res = square(tf.constant([1, 2, 3, 4, 5]))
print(res)
print(square.graph.get_operations()[:3])
Результат выполнения кода:
tf.Tensor([ 1 4 9 16 25], shape=(5,), dtype=int32)
[<tf.Operation 't' type=Placeholder>, <tf.Operation 'mul' type=Mul>]
Смотрите также
-
функцию
print,
которая выводит значения внутри графа -
функцию
custom_gradient,
которая задаёт собственный градиент функции -
функцию
device,
которая указывает устройство для операций -
функцию
ensure_shape,
которая проверяет форму тензора внутри графа