tf.py_function
| Просмотреть исходный код на GitHub |
Оборачивает функцию Python в операцию TensorFlow, которая выполняет её жадно.
tf.py_function(
func, inp, Tout, name=None
)
Эта функция позволяет выражать вычисления в графе TensorFlow как функции Python. В частности, она оборачивает функцию Python func в один раз дифференцируемую операцию TensorFlow, которая выполняет её с включённым жадным выполнением. Вследствие этого, tf.py_function позволяет выражать управляющие конструкции с помощью конструкций Python (if, while, for, и т. д.), а не конструкций управления потоком TensorFlow (tf.cond, tf.while_loop). Например, вы можете использовать tf.py_function для реализации функции log-huber:
def log_huber(x, m):
if tf.abs(x) <= m:
return x**2
else:
return m**2 * (1 - 2 * tf.math.log(m) + tf.math.log(x**2))
x = tf.constant(1.0)
m = tf.constant(2.0)
with tf.GradientTape() as t:
t.watch([x, m])
y = tf.py_function(func=log_huber, inp=[x, m], Tout=tf.float32)
dy_dx = t.gradient(y, x)
assert dy_dx.numpy() == 2.0
Вы также можете использовать tf.py_function для отладки своих моделей во время выполнения с помощью инструментов Python, т. е. вы можете изолировать части вашего кода, которые вы хотите отладить, обернуть их в функции Python и вставить pdb точки отслеживания или операторы вывода по своему усмотрению, и обернуть эти функции в tf.py_function.
Для получения дополнительной информации о жадном выполнении см. Руководство по жадному выполнению.
tf.py_function по духу похожа на tf.compat.v1.py_func, но в отличие от последней, первая позволяет использовать операции TensorFlow в обернутой функции Python. В частности, в то время как tf.compat.v1.py_func выполняется только на ЦП и оборачивает функции, которые принимают массивы NumPy в качестве входных данных и возвращают массивы NumPy в качестве выходных данных, tf.py_function может размещаться на графических процессорах и оборачивает функции, которые принимают тензоры в качестве входных данных, выполняют операции TensorFlow в своих телах и возвращают тензоры в качестве выходных данных.
Примечание: Мы рекомендуем избегать использования tf.py_function за пределами прототипирования и экспериментов из-за следующих известных ограничений:
Вызов
tf.py_functionполучит блокировку интерпретатора Python (GIL), которая позволяет одному потоку выполняться в любой момент времени. Это исключит эффективное распараллеливание и распределение выполнения программы.Тело функции (т. е.
func) не будет сериализовано вGraphDef. Поэтому не следует использовать эту функцию, если вам нужно сериализовать свою модель и восстановить её в другой среде.Операция должна выполняться в том же адресном пространстве, что и программа Python, вызывающая
tf.py_function(). Если вы используете распределённый TensorFlow, вы должны запуститьtf.distribute.Serverв том же процессе, что и программа, вызывающаяtf.py_function(), и вы должны привязать созданную операцию к устройству в этом сервере (например, с помощьюwith tf.device():).В настоящее время
tf.py_functionнесовместим с XLA. Вызовtf.py_functionвнутриtf.function(jit_comiple=True)вызовет ошибку.
| Аргументы | |
|---|---|
func | Функция Python, принимающая inp в качестве аргументов и возвращающая значение (или список значений), тип которого описан в Tout. |
inp | Входные аргументы для func. Список, элементы которого являются Tensor или CompositeTensors (например, tf.RaggedTensor); или один Tensor или CompositeTensor. |
Tout | Тип(ы) возвращаемого(ых) значения(й) func. Один из следующих:
|
name | Имя операции (необязательно). |
| Возвращаемое значение | |
|---|---|
Значение(я), вычисленное(ые) func: Tensor, CompositeTensor, или список Tensor и CompositeTensor; или пустой список, если func возвращает None. |
© 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/py_function