Melhor desempenho com tf.function

Veja no TensorFlow.org Executar no Google Colab Ver fonte no GitHubBaixar caderno

No TensorFlow 2, a execução antecipada é ativada por padrão. A interface do usuário é intuitiva e flexível (executar operações pontuais é muito mais fácil e rápido), mas isso pode prejudicar o desempenho e a capacidade de implantação.

Você pode usar tf.function para fazer gráficos de seus programas. É uma ferramenta de transformação que cria gráficos de fluxo de dados independentes de Python a partir de seu código Python. Isso ajudará você a criar modelos portáteis e de alto desempenho, e é necessário usar SavedModel .

Este guia irá ajudá-lo a conceituar como o tf.function funciona nos bastidores, para que você possa usá-lo de forma eficaz.

As principais dicas e recomendações são:

  • Depure no modo ansioso e depois decore com @tf.function .
  • Não confie nos efeitos colaterais do Python, como mutação de objeto ou anexos de lista.
  • tf.function funciona melhor com operações do TensorFlow; As chamadas NumPy e Python são convertidas em constantes.

Configurar

import tensorflow as tf

Defina uma função auxiliar para demonstrar os tipos de erros que você pode encontrar:

import traceback
import contextlib

# Some helper code to demonstrate the kinds of errors you might encounter.
@contextlib.contextmanager
def assert_raises(error_class):
  try:
    yield
  except error_class as e:
    print('Caught expected exception \n  {}:'.format(error_class))
    traceback.print_exc(limit=2)
  except Exception as e:
    raise e
  else:
    raise Exception('Expected {} to be raised but no error was raised!'.format(
        error_class))

Fundamentos

Uso

Uma Function que você define (por exemplo, aplicando o decorador @tf.function ) é como uma operação principal do TensorFlow: você pode executá-la rapidamente; você pode calcular gradientes; e assim por diante.

@tf.function  # The decorator converts `add` into a `Function`.
def add(a, b):
  return a + b

add(tf.ones([2, 2]), tf.ones([2, 2]))  #  [[2., 2.], [2., 2.]]
<tf.Tensor: shape=(2, 2), dtype=float32, numpy=
array([[2., 2.],
       [2., 2.]], dtype=float32)>
v = tf.Variable(1.0)
with tf.GradientTape() as tape:
  result = add(v, 1.0)
tape.gradient(result, v)
<tf.Tensor: shape=(), dtype=float32, numpy=1.0>

Você pode usar Function dentro de outras Function .

@tf.function
def dense_layer(x, w, b):
  return add(tf.matmul(x, w), b)

dense_layer(tf.ones([3, 2]), tf.ones([2, 2]), tf.ones([2]))
<tf.Tensor: shape=(3, 2), dtype=float32, numpy=
array([[3., 3.],
       [3., 3.],
       [3., 3.]], dtype=float32)>

Function s pode ser mais rápida que o código ansioso, especialmente para gráficos com muitas operações pequenas. Mas para gráficos com algumas operações caras (como convoluções), você pode não ver muita aceleração.

import timeit
conv_layer = tf.keras.layers.Conv2D(100, 3)

@tf.function
def conv_fn(image):
  return conv_layer(image)

image = tf.zeros([1, 200, 200, 100])
# Warm up
conv_layer(image); conv_fn(image)
print("Eager conv:", timeit.timeit(lambda: conv_layer(image), number=10))
print("Function conv:", timeit.timeit(lambda: conv_fn(image), number=10))
print("Note how there's not much difference in performance for convolutions")
Eager conv: 0.006058974999177735
Function conv: 0.005791576000774512
Note how there's not much difference in performance for convolutions

Rastreamento

Esta seção expõe como o Function funciona nos bastidores, incluindo detalhes de implementação que podem mudar no futuro . No entanto, uma vez que você entenda por que e quando o rastreamento acontece, é muito mais fácil usar o tf.function efetivamente!

O que é "rastreamento"?

Uma Function executa seu programa em um gráfico do TensorFlow . No entanto, um tf.Graph não pode representar todas as coisas que você escreveria em um programa TensorFlow ansioso. Por exemplo, Python suporta polimorfismo, mas tf.Graph requer que suas entradas tenham um tipo de dados e uma dimensão especificados. Ou você pode realizar tarefas secundárias como ler argumentos de linha de comando, gerar um erro ou trabalhar com um objeto Python mais complexo; nenhuma dessas coisas pode ser executada em um tf.Graph .

Function preenche essa lacuna separando seu código em dois estágios:

1) Na primeira etapa, chamada de " tracing ", Function cria um novo tf.Graph . O código Python é executado normalmente, mas todas as operações do TensorFlow (como adicionar dois tensores) são adiadas : elas são capturadas pelo tf.Graph e não são executadas.

2) Na segunda etapa, é executado um tf.Graph que contém tudo o que foi adiado na primeira etapa. Este estágio é muito mais rápido que o estágio de rastreamento.

Dependendo de suas entradas, Function nem sempre executará o primeiro estágio quando for chamada. Consulte "Regras de rastreamento" abaixo para ter uma noção melhor de como ele faz essa determinação. Ignorar o primeiro estágio e executar apenas o segundo estágio é o que oferece o alto desempenho do TensorFlow.

Quando Function decide rastrear, o estágio de rastreamento é imediatamente seguido pelo segundo estágio, portanto, chamar a Function cria e executa o tf.Graph . Mais tarde, você verá como pode executar apenas o estágio de rastreamento com get_concrete_function .

Quando você passa argumentos de tipos diferentes para um Function , ambos os estágios são executados:

@tf.function
def double(a):
  print("Tracing with", a)
  return a + a

print(double(tf.constant(1)))
print()
print(double(tf.constant(1.1)))
print()
print(double(tf.constant("a")))
print()
Tracing with Tensor("a:0", shape=(), dtype=int32)
tf.Tensor(2, shape=(), dtype=int32)

Tracing with Tensor("a:0", shape=(), dtype=float32)
tf.Tensor(2.2, shape=(), dtype=float32)

Tracing with Tensor("a:0", shape=(), dtype=string)
tf.Tensor(b'aa', shape=(), dtype=string)

Observe que, se você chamar repetidamente uma Function com o mesmo tipo de argumento, o TensorFlow pulará o estágio de rastreamento e reutilizará um gráfico rastreado anteriormente, pois o gráfico gerado seria idêntico.

# This doesn't print 'Tracing with ...'
print(double(tf.constant("b")))
tf.Tensor(b'bb', shape=(), dtype=string)

Você pode usar pretty_printed_concrete_signatures() para ver todos os traços disponíveis:

print(double.pretty_printed_concrete_signatures())
double(a)
  Args:
    a: int32 Tensor, shape=()
  Returns:
    int32 Tensor, shape=()

double(a)
  Args:
    a: float32 Tensor, shape=()
  Returns:
    float32 Tensor, shape=()

double(a)
  Args:
    a: string Tensor, shape=()
  Returns:
    string Tensor, shape=()

Até agora, você viu que tf.function cria uma camada de despacho dinâmica em cache sobre a lógica de rastreamento de gráfico do TensorFlow. Para ser mais específico sobre a terminologia:

  • Um tf.Graph é a representação bruta, independente de linguagem e portátil de uma computação do TensorFlow.
  • Um ConcreteFunction envolve um tf.Graph .
  • Uma Function gerencia um cache de ConcreteFunction se escolhe o caminho certo para suas entradas.
  • tf.function envolve uma função Python, retornando um objeto Function .
  • O rastreamento cria um tf.Graph e o envolve em um ConcreteFunction , também conhecido como rastreamento.

Regras de rastreamento

Uma Function determina se deve ser reutilizada uma ConcreteFunction rastreada calculando uma chave de cache dos argumentos e kwargs de uma entrada. Uma chave de cache é uma chave que identifica uma ConcreteFunction com base nos argumentos e kwargs de entrada da chamada da