Spec-Zone.ru › TensorFlow

tf.types.experimental.TraceType

Представляет тип объекта(ов) для целей отслеживания tf.function.

Использование в блокнотах

Используется в руководстве
  • Улучшенная производительность с 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):
    self.fruit_type = type(fruit)
    self.fruit_value = fruit

  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

  def placeholder_value(self, placeholder_context=None):
    return self.fruit_value

class Fruit:

 def __tf_tracing_type__(self, context):
   return FruitTraceType(self)

Теперь, если мы попробуем вызвать его снова:

get_mixed_flavor(Apple(), Mango()) # Traces a new concrete function
get_mixed_flavor(Apple(), Mango()) # Re-uses the traced concrete function

Методы

cast

Просмотреть исходный код

cast(
    value, cast_context
) -> Any

Приведение значения к этому типу.

Аргументы
value Входное значение, принадлежащее этому TraceType.
cast_context Контекст, предназначенный для внутреннего/будущего использования.
Возвращает
Значение, приведённое к этому TraceType.
Исключения
AssertionError Когда _cast не перегружен в подклассе, значение возвращается непосредственно, и оно должно быть таким же, как self.placeholder_value().

flatten

Просмотреть исходный код

flatten() -> List['TraceType']

Возвращает список TensorSpec, соответствующих значениям to_tensors.

from_tensors

Просмотреть исходный код

from_tensors(
    tensors: Iterator[core.Tensor]
) -> Any

Генерирует значение этого типа из тензоров.

Должен использовать то же фиксированное количество тензоров, что и to_tensors.

Аргументы
tensors Итератор, из которого можно извлечь тензоры.
Возвращает
Значение этого типа.

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 Последовательность TraceTypes.

Пример:

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)

placeholder_value

Просмотреть исходный код

@abc.abstractmethod
placeholder_value(
    placeholder_context
) -> Any

Создаёт заполнитель для отслеживания.

tf.function отслеживает с помощью значения заполнителя, а не фактического значения. Например, значение заполнителя может представлять собой несколько разных фактических значений. Это означает, что отслеживание, сгенерированное с помощью этого значения заполнителя, является более общим и повторно используемым, что экономит дорогостоящее повторное отслеживание.

Аргументы
placeholder_context Контекст, предназначенный для внутреннего/будущего использования.

Для примера Fruit выше, реализация:

class FruitTraceType:
  def placeholder_value(self, placeholder_context):
    return Fruit()

указывает tf.function отслеживать с объектами Fruit() вместо фактических объектов Apple() и Mango(), когда он получает вызов get_mixed_flavor(Apple(), Mango()). Например, аргументы Tensor заменяются на Tensors с аналогичной формой и типом данных, выводом из операции tf.Placeholder.

В более общем плане, значения заполнителей являются аргументами tf.function, как видно из тела функции:

@tf.function
def foo(x):
  # Here `x` is be the placeholder value
  ...

foo(x) # Here `x` is the actual value

to_tensors

Просмотреть исходный код

to_tensors(
    value: Any
) -> List[core.Tensor]

Разбивает значение этого типа на тензоры.

Для TraceType количество сгенерированных тензоров для соответствующего значения должно быть постоянным.

Аргументы
value Значение, принадлежащее этому TraceType
Возвращает
Список тензоров.

__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/api_docs/python/tf/types/experimental/TraceType

Spec-Zone.ru

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