tf.function
| Просмотреть исходный код на GitHub |
Компилирует функцию в вызываемый граф TensorFlow. (устаревшие аргументы) (устаревшие аргументы)
tf.function(
func=None,
input_signature=None,
autograph=True,
jit_compile=None,
reduce_retracing=False,
experimental_implements=None,
experimental_autograph_options=None,
experimental_relax_shapes=None,
experimental_compile=None,
experimental_follow_type_hints=None
) -> tf.types.experimental.GenericFunction
tf.function создаёт tf.types.experimental.GenericFunction, который выполняет граф TensorFlow (tf.Graph), созданный путем трассировки компиляции операций TensorFlow в func. Дополнительную информацию по этой теме можно найти в Введении в графы и tf.function.
См. Улучшение производительности с помощью tf.function для советов по производительности и известных ограничений.
Пример использования:
@tf.function def f(x, y): return x ** 2 + y x = tf.constant([2, 3]) y = tf.constant([3, -2]) f(x, y) <tf.Tensor: ... numpy=array([7, 7], ...)>
Трассировка компиляции позволяет выполнять операции, не относящиеся к TensorFlow, но при особых условиях. В общем случае гарантируется выполнение только операций TensorFlow и создание новых результатов каждый раз при вызове GenericFunction.
Особенности
func может использовать операторы Python, зависящие от данных, включая if, for, while break, continue и return:
@tf.function
def f(x):
if tf.reduce_sum(x) > 0:
return x * x
else:
return -x // 2
f(tf.constant(-2))
<tf.Tensor: ... numpy=1>
Закрытие func может включать объекты tf.Tensor и tf.Variable:
@tf.function def f(): return x ** 2 + y x = tf.constant([-2, -3]) y = tf.Variable([3, -2]) f() <tf.Tensor: ... numpy=array([7, 7], ...)>
func также может использовать операции с побочными эффектами, такие как tf.print, tf.Variable и другие:
v = tf.Variable(1)
@tf.function
def f(x):
for i in tf.range(x):
v.assign_add(i)
f(3)
v
<tf.Variable ... numpy=4>
l = []
@tf.function
def f(x):
for i in x:
l.append(i + 1) # Caution! Will only happen once when tracing
f(tf.constant([1, 2, 3]))
l
[<tf.Tensor ...>]
Вместо этого используйте коллекции TensorFlow, такие как tf.TensorArray:
@tf.function
def f(x):
ta = tf.TensorArray(dtype=tf.int32, size=0, dynamic_size=True)
for i in range(len(x)):
ta = ta.write(i, x[i] + 1)
return ta.stack()
f(tf.constant([1, 2, 3]))
<tf.Tensor: ..., numpy=array([2, 3, 4], ...)>
tf.function создает полиморфные вызываемые функции
tf.types.experimental.GenericFunction может содержать несколько tf.types.experimental.ConcreteFunction внутри, каждая из которых специализирована для аргументов с разными типами данных или формами, так как TensorFlow может выполнять больше оптимизаций в графах с определёнными формами, типами данных и значениями постоянных аргументов. tf.function рассматривает любые чистые значения Python как нечёткие объекты (лучше всего представляемые как константы времени компиляции) и создаёт отдельный tf.Graph для каждого набора аргументов Python, с которыми он сталкивается. Дополнительную информацию см. в руководстве по tf.function.
Выполнение GenericFunction выберет и выполнит соответствующее ConcreteFunction на основе типов и значений аргументов.
Чтобы получить отдельное ConcreteFunction, используйте метод GenericFunction.get_concrete_function. Его можно вызвать с теми же аргументами, что и func, и он возвращает tf.types.experimental.ConcreteFunction. ConcreteFunction основаны на единственном tf.Graph:
@tf.function def f(x): return x + 1 isinstance(f.get_concrete_function(1).graph, tf.Graph) True
ConcreteFunction можно выполнять так же, как и GenericFunction, но их вход ограничен типами, для которых они специализированы.
Перестроение
ConcreteFunctions создаются (отслеживаются) на лету, поскольку GenericFunction вызывается с новыми типами или формами TensorFlow или новыми значениями Python в качестве аргументов. Когда GenericFunction создает новую трассировку, говорят, что func перестраивается. Перестроение часто является проблемой производительности для tf.function, так как оно может быть значительно медленнее, чем выполнение уже отслеженного графа. Желательно минимизировать количество перестроек в вашем коде.
@tf.function def f(x): return tf.abs(x) f1 = f.get_concrete_function(1) f2 = f.get_concrete_function(2) # Slow - compiles new graph f1 is f2 False f1 = f.get_concrete_function(tf.constant(1)) f2 = f.get_concrete_function(tf.constant(2)) # Fast - reuses f1 f1 is f2 True
Численные аргументы Python следует использовать только тогда, когда они принимают несколько различных значений, таких как гиперпараметры, например, количество слоёв в нейронной сети.
Подписи ввода
Для тензорных аргументов, GenericFunction создаёт новую ConcreteFunction для каждого уникального набора форм и типов входных данных. Пример ниже создаёт две отдельные ConcreteFunction , каждая из которых специализируется на другой форме:
@tf.function def f(x): return x + 1 vector = tf.constant([1.0, 1.0]) matrix = tf.constant([[3.0]]) f.get_concrete_function(vector) is f.get_concrete_function(matrix) False
«Подпись ввода» может быть необязательно предоставлена tf.function для управления этим процессом. Подпись ввода определяет форму и тип каждого тензорного аргумента функции с помощью объекта tf.TensorSpec. Можно использовать более общие формы. Это гарантирует создание только одного ConcreteFunction и ограничивает GenericFunction указанными формами и типами. Это эффективный способ ограничить перестроение, когда у тензоров динамические формы.
@tf.function(
input_signature=[tf.TensorSpec(shape=None, dtype=tf.float32)])
def f(x):
return x + 1
vector = tf.constant([1.0, 1.0])
matrix = tf.constant([[3.0]])
f.get_concrete_function(vector) is f.get_concrete_function(matrix)
True
Переменные могут быть созданы только один раз
tf.function разрешает создание новых объектов tf.Variable только при первом вызове:
class MyModule(tf.Module):
def __init__(self):
self.v = None
@tf.function
def __call__(self, x):
if self.v is None:
self.v = tf.Variable(tf.ones_like(x))
return self.v * x
В целом рекомендуется создавать tf.Variable вне tf.function. В простых случаях состояние, сохраняющееся через границы tf.function, может быть реализовано с помощью чисто функционального стиля, в котором состояние представлено tf.Tensor , передаваемыми в качестве аргументов и возвращаемыми в качестве возвращаемых значений.
Сравните два стиля ниже:
state = tf.Variable(1) @tf.function def f(x): state.assign_add(x) f(tf.constant(2)) # Non-pure functional style state <tf.Variable ... numpy=3>
state = tf.constant(1) @tf.function def f(state, x): state += x return state state = f(state, tf.constant(2)) # Pure functional style state <tf.Tensor: ... numpy=3>
Операции Python выполняются только один раз за трассировку
func может содержать операции TensorFlow, смешанные с чистыми операциями Python. Однако при выполнении функции будут выполняться только операции TensorFlow. Операции Python выполняются только один раз — во время трассировки. Если операции TensorFlow зависят от результатов операций Python, эти результаты будут заморожены в графе.
@tf.function
def f(a, b):
print('this runs at trace time; a is', a, 'and b is', b)
return b
f(1, tf.constant(1))
this runs at trace time; a is 1 and b is Tensor("...", shape=(), dtype=int32)
<tf.Tensor: shape=(), dtype=int32, numpy=1>
f(1, tf.constant(2)) <tf.Tensor: shape=(), dtype=int32, numpy=2>
f(2, tf.constant(1))
this runs at trace time; a is 2 and b is Tensor("...", shape=(), dtype=int32)
<tf.Tensor: shape=(), dtype=int32, numpy=1>
f(2, tf.constant(2)) <tf.Tensor: shape=(), dtype=int32, numpy=2>
Использование аннотаций типов для повышения производительности
'experimental_follow_type_hints` можно использовать вместе с аннотациями типов для уменьшения перестроения, автоматически преобразуя любые значения Python в tf.Tensor (это не делается по умолчанию, если вы не используете подписи ввода).
@tf.function(experimental_follow_type_hints=True)
def f_with_hints(x: tf.Tensor):
print('Tracing')
return x
@tf.function(experimental_follow_type_hints=False)
def f_no_hints(x: tf.Tensor):
print('Tracing')
return x
f_no_hints(1)
Tracing
<tf.Tensor: shape=(), dtype=int32, numpy=1>
f_no_hints(2)
Tracing
<tf.Tensor: shape=(), dtype=int32, numpy=2>
f_with_hints(1)
Tracing
<tf.Tensor: shape=(), dtype=int32, numpy=1>
f_with_hints(2)
<tf.Tensor: shape=(), dtype=int32, numpy=2>
| Аргументы | |
|---|---|
func | функция, подлежащая компиляции. Если func равно None, tf.function возвращает декоратор, который можно вызвать с одним аргументом — func. Другими словами, tf.function(input_signature=...)(func) эквивалентно tf.function(func, input_signature=...). Первый можно использовать как декоратор. |
input_signature | Возможная вложенная последовательность объектов tf.TensorSpec, задающих формы и типы данных тензоров, которые будут переданы в эту функцию. Если None, для каждой выведенной подписи входных данных создается отдельная функция. Если input_signature задано, каждый вход в func должен быть Tensor, и func не может принимать **kwargs. |
autograph | Необходимо ли применять автографирование к func перед трассировкой графа. Для инструкций Python с условными операторами, зависящими от данных, требуется autograph=True. Дополнительная информация в руководстве tf.function и AutoGraph. |
jit_compile | Если True, компилирует функцию с использованием XLA. XLA выполняет оптимизации компилятора, такие как слияние, и пытается сгенерировать более эффективный код. Это может значительно улучшить производительность. Если установлено True, вся функция должна быть компилируемой с помощью XLA, в противном случае выбрасывается errors.InvalidArgumentError. Если None (значение по умолчанию), компилирует функцию с XLA при выполнении на TPU и использует обычный путь выполнения функции при выполнении на других устройствах. Если False, выполняет функцию без компиляции с XLA. Установите это значение в False при непосредственном запуске функции на нескольких устройствах на TPU (например, двух ядрах TPU, одном ядре TPU и его процессоре хоста). Не все функции компилируемы, см. список острых углов. |
reduce_retracing | Если True, tf.function пытается уменьшить объем повторной трассировки, например, используя более общие формы. Это можно контролировать для пользовательских объектов, настраивая связанные tf.types.experimental.TraceType. |
experimental_implements | При необходимости содержит имя «известной» функции, которую она реализует. Например, "mycompany.my_recurrent_cell". Это хранится в качестве атрибута в функции вывода, который затем можно обнаружить при обработке сериализованной функции. См. стандартизацию составных операций для получения подробной информации. Пример использования этого атрибута приведен в примере. Приведенный выше код автоматически обнаруживает и заменяет функцию, реализующую «embedded_matmul», что позволяет TFLite заменить свои собственные реализации. Например, пользователь TensorFlow может использовать этот атрибут, чтобы отметить, что его функция также реализует embedded_matmul (возможно, более эффективно!), указав его с помощью этого параметра: @tf.function(experimental_implements="embedded_matmul"). Это можно указать как просто строковое имя функции, или как NameAttrList, соответствующее списку пар «ключ-значение», связанных с именем функции. Имя функции будет находиться в поле 'name' объекта NameAttrList. Для определения формальной операции TF для этой функции обратитесь к экспериментальному проекту составного TF. |
experimental_autograph_options | Необязательная кортеж значений tf.autograph.experimental.Feature. |
experimental_relax_shapes | Устарело. Используйте reduce_retracing вместо этого. |
experimental_compile | Устаревшее псевдоним для 'jit_compile'. |
experimental_follow_type_hints | Если True, функция может использовать аннотации типов из func для оптимизации производительности трассировки. Например, аргументы, аннотированные tf.Tensor, будут автоматически преобразованы в тензор. |
| Возвращаемое значение | |
|---|---|
Если func не равно None, возвращает tf.types.experimental.GenericFunction. Если func равно None, возвращает декоратор, который при вызове с одним аргументом func возвращает tf.types.experimental.GenericFunction. |
| Исключения | |
|---|---|
ValueError при попытке использовать jit_compile=True, но поддержка XLA недоступна. |
© 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/function