tf.types.experimental.PolymorphicFunction
Базовый класс для полиморфных функций графа.
Наследуется от: Callable
Функции графа — это объекты Python, вызываемые для перенаправления вызовов в граф TensorFlow. Полиморфные функции графа могут быть основаны на нескольких графах TF и автоматически выбирают соответствующую специализацию на основе типа входных данных, с которыми они вызывались. Они также могут создавать специализации на лету, например, путем отслеживания.
Также см. tf.function.
| Атрибуты | |
|---|---|
function_type | Возвращает FunctionType, описывающий этот вызываемый объект. |
Методы
experimental_get_compiler_ir
experimental_get_compiler_ir(
*args, **kwargs
)
Возвращает IR компилятора для скомпилированной функции.
Этот API предназначен только для отладки, так как нет гарантий обратной совместимости возвращаемого IR или допустимых значений stage.
| Аргументы | |
|---|---|
*args | аргументы компиляции поддерживают входы: (1) все входы — TensorSpec или (2) все входы — tf.Tensor/переменные Python. |
**kwargs | Параметры, используемые для компиляции. Требования такие же, как и для аргументов компиляции. |
| Возвращаемое значение |
|---|
Функция, вызываемая со следующими ключевыми аргументами:
Например, для @tf.function(jit_compile=True) def f(x): return x + 1 f.experimental_get_compiler_ir(tf.random.normal([10, 10])(stage='hlo') выход будет: HloModule a_inference_f_13__.9
ENTRY %a_inference_f_13__.9 (arg0.1: f32[10,10]) -> f32[10,10] {
%arg0.1 = f32[10,10]{1,0} parameter(0), parameter_replication={false}
%reshape.2 = f32[10,10]{1,0} reshape(f32[10,10]{1,0} %arg0.1)
%constant.3 = f32[] constant(1)
%broadcast.4 = f32[10,10]{1,0} broadcast(f32[] %constant.3)
%add.5 = f32[10,10]{1,0} add(f32[10,10]{1,0} %reshape.2,
f32[10,10]{1,0} %broadcast.4)
%reshape.6 = f32[10,10]{1,0} reshape(f32[10,10]{1,0} %add.5)
%tuple.7 = (f32[10,10]{1,0}) tuple(f32[10,10]{1,0} %reshape.6)
ROOT %get-tuple-element.8 = f32[10,10]{1,0}
get-tuple-element((f32[10,10]{1,0}) %tuple.7), index=0
}
Вот еще один пример с использованием tf.TensorSpec в качестве входных данных: y = tf.Variable(tf.zeros([10, 20], dtype=tf.float32)) @tf.function(jit_compile=True) def f(x): return x + y hlo_str = f.experimental_get_compiler_ir(tf.TensorSpec(shape=(10, 20)))(stage='hlo') Выход: HloModule a_inference_f_120__.8,
entry_computation_layout={(f32[10,20]{1,0},f32[10,20]{1,0})->f32[10,20]{1,0} }
ENTRY %a_inference_f_120__.8 (arg0.1: f32[10,20], arg1.2: f32[10,20]) ->
f32[10,20] {
%arg0.1 = f32[10,20]{1,0} parameter(0), parameter_replication={false},
metadata={op_name="XLA_Args"}
%reshape.3 = f32[10,20]{1,0} reshape(f32[10,20]{1,0} %arg0.1)
%arg1.2 = f32[10,20]{1,0} parameter(1), parameter_replication={false},
metadata={op_name="XLA_Args"}
%add.4 = f32[10,20]{1,0} add(f32[10,20]{1,0} %reshape.3, f32[10,20]{1,0}
%arg1.2), metadata={op_type="AddV2" op_name="add"
source_file="<ipython-input-16-ea04879c1873>" source_line=4}
%reshape.5 = f32[10,20]{1,0} reshape(f32[10,20]{1,0} %add.4),
metadata={op_name="XLA_Retvals"}
%tuple.6 = (f32[10,20]{1,0}) tuple(f32[10,20]{1,0} %reshape.5),
metadata={op_name="XLA_Retvals"}
ROOT %get-tuple-element.7 = f32[10,20]{1,0}
get-tuple-element((f32[10,20]{1,0}) %tuple.6), index=0,
metadata={op_name="XLA_Retvals"}
}
</td>
</tr>
</table>
Модуль HLO принимает плоский список входных данных. Чтобы получить порядок подписей этих входных данных, пользователи могут вызвать # Use concrete_fn to get the hlo_module flat_args.
concrete_fn = f.get_concrete_function(tf.TensorSpec(shape=(10, 20)))
flat_args = list(
tf.nest.flatten(concrete_fn.structured_input_signature)
) + concrete_fn.captured_inputs
|
|||||||||||||||||||||||||||
| Аргументы | |
|---|---|
*args | входные данные для специализации. |
**kwargs | входные данные для специализации. |
| Возвращаемое значение | |
|---|---|
ConcreteFunction. |
__call__
__call__(
*args, **kwargs
)
Выполняет этот вызываемый объект.
Это ведет себя как обычная операция - в режиме eager она немедленно запускает выполнение, возвращая результаты. В режиме графа она создает операции, которые возвращают символические значения TensorFlow (например, tf.Tensor, tf.data.Dataset и т. д.). Например, вызываемые объекты tf.function обычно генерируют операцию tf.raw_ops.PartitionedCall, но не всегда - точные генерируемые операции являются внутренней реализацией.
| Аргументы | |
|---|---|
*args | позиционный аргумент для этого вызова |
**kwargs | ключевые аргументы для этого вызова |
| Возвращаемое значение | |
|---|---|
| Результаты выполнения. |
За исключением случаев, когда указано иное, контент этой страницы лицензирован по лицензии Creative Commons Attribution 4.0, а примеры кода лицензированы по лицензии Apache 2.0. Подробнее см. Политику Google для разработчиков. Java — зарегистрированный товарный знак Oracle и/или ее дочерних компаний. Некоторые материалы лицензированы по лицензии numpy.
Последнее обновление 2024-04-26 UTC.
© 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/api_docs/python/tf/types/experimental/PolymorphicFunction