tf.types.experimental.TraceType
Представляет тип объекта(ов) для целей отслеживания tf.function.
TraceType — это абстрактный класс, от которого другие классы могут наследоваться, чтобы предоставить информацию о связанных классах для целей отслеживания tf.function. Логика типизации, предоставляемая этим механизмом, будет использоваться для принятия решений о использовании кэшированных конкретных функций и повторном отслеживании.
Например, если у нас есть следующие tf.function и классы:
@tf.function def get_mixed_flavor(fruit_a, fruit_b): return fruit_a.flavor + fruit_b.flavor class Fruit: flavor = tf.constant([0, 0]) class Apple(Fruit): flavor = tf.constant([1, 2]) class Mango(Fruit): flavor = tf.constant([3, 4])
tf.function не знает, когда следует повторно использовать существующую конкретную функцию в отношении класса Fruit, поэтому по умолчанию он повторно отслеживает для каждого нового экземпляра.
get_mixed_flavor(Apple(), Mango()) # Traces a new concrete function get_mixed_flavor(Apple(), Mango()) # Traces a new concrete function again
Однако мы, как разработчики класса Fruit, знаем, что каждый подкласс имеет фиксированный вид, и мы можем повторно использовать существующую отслеженную конкретную функцию, если это был тот же подкласс. Избежание такого ненужного отслеживания конкретных функций может принести значительные преимущества производительности.
class FruitTraceType(tf.types.experimental.TraceType):
def __init__(self, fruit_type):
self.fruit_type = fruit_type
def is_subtype_of(self, other):
return (type(other) is FruitTraceType and
self.fruit_type is other.fruit_type)
def most_specific_common_supertype(self, others):
return self if all(self == other for other in others) else None
class Fruit:
def __tf_tracing_type__(self, context):
return FruitTraceType(type(self))
Теперь, если мы попробуем вызвать его снова:
get_mixed_flavor(Apple(), Mango()) # Traces a new concrete function get_mixed_flavor(Apple(), Mango()) # Re-uses the traced concrete function
Методы
is_subtype_of
@abc.abstractmethod
is_subtype_of(
other: 'TraceType'
) -> bool
Возвращает True, если self является подтипом other.
Например, tf.function использует подтипизацию для диспетчеризации: если a.is_subtype_of(b) имеет значение True, то аргумент типа TraceType a может быть использован в качестве аргумента для ConcreteFunction отслеженного с помощью TraceType b.
| Аргументы | |
|---|---|
other | Объект TraceType для сравнения. |
Пример:
class Dimension(TraceType):
def __init__(self, value: Optional[int]):
self.value = value
def is_subtype_of(self, other):
# Either the value is the same or other has a generalized value that
# can represent any specific ones.
return (self.value == other.value) or (other.value is None)
most_specific_common_supertype
@abc.abstractmethod
most_specific_common_supertype(
others: Sequence['TraceType']
) -> Optional['TraceType']
Возвращает наиболее специфичный супертип self и others, если он существует.
Возвращаемый TraceType является супертипом self и others, то есть они все являются подтипами (см. is_subtype_of) его. Он также наиболее специфичен, то есть у него нет подтипа, который также является общим супертипом self и others.
Если self и others не имеют общего супертипа, возвращается None.
| Аргументы | |
|---|---|
others | Последовательность объектов TraceType. |
Пример:
class Dimension(TraceType):
def __init__(self, value: Optional[int]):
self.value = value
def most_specific_common_supertype(self, other):
# Either the value is the same or other has a generalized value that
# can represent any specific ones.
if self.value == other.value:
return self.value
else:
return Dimension(None)
__eq__
@abc.abstractmethod
__eq__(
other
) -> bool
Возвращает self==value.
© 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/types/experimental/TraceType