Spec-Zone.ru › TensorFlow 2.9

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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API