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):
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