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.compat.v1.placeholder(tf.float32)
m = tf.compat.v1.placeholder(tf.float32)
y = tf.py_function(func=log_huber, inp=[x, m], Tout=tf.float32)
dy_dx = tf.gradients(y, x)[0]
with tf.compat.v1.Session() as sess:
# The session executes `log_huber` eagerly. Given the feed values below,
# it will take the first branch, so `y` evaluates to 1.0 and
# `dy_dx` evaluates to 2.0.
y, dy_dx = sess.run([y, dy_dx], feed_dict={x: 1.0, m: 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 может быть размещена на GPU и оборачивает функции, которые принимают тензоры в качестве входных данных, выполняют операции TensorFlow в своих телах и возвращают тензоры в качестве выходных данных.
Как и tf.compat.v1.py_func, tf.py_function имеет следующие ограничения в отношении сериализации и распределения:
Тело функции (т.е.
func) не будет сериализовано вGraphDef. Поэтому не следует использовать эту функцию, если вам необходимо сериализовать модель и восстановить её в другой среде.Операция должна выполняться в том же адресном пространстве, что и программа Python, которая вызывает
tf.py_function(). Если вы используете распределённый TensorFlow, вы должны запуститьtf.distribute.Serverв том же процессе, что и программа, которая вызываетtf.py_function(), и вы должны привязать созданную операцию к устройству в этом сервере (например, с помощьюwith tf.device():).
| Аргументы | |
|---|---|
func | Функция Python, которая принимает список объектов Tensor с типами элементов, соответствующими соответствующим объектам tf.Tensor в inp и возвращает список объектов Tensor (или один Tensor, или None) с типами элементов, соответствующими соответствующим значениям в Tout. |
inp | Список объектов Tensor . |
Tout | Список или кортеж типов данных TensorFlow или один тип данных TensorFlow, если их только один, указывающий, что возвращает func; пустой список, если не возвращается значение (т.е., если возвращаемое значение None). |
name | Имя операции (необязательно). |
| Возвращаемое значение | |
|---|---|
Список объектов Tensor или один Tensor, который func вычисляет; пустой список, если func возвращает None. |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.3/api_docs/python/tf/py_function