Salvamento e Carregamento de Modelos no TensorFlow

O TensorFlow oferece dois métodos principais para implementar o salvamento e carregamento de modelos: a abordagem via tf.keras e a abordagem nativa do tf.train.

1. Salvamento e Carregamento com tf.keras

O módulo tf.keras do TensorFlow fornece uma abordagem simples e robusta para gerenciar modelos.

import tensorflow as tf

# Construção de uma rede neural básica
rede_neural = tf.keras.Sequential([
    tf.keras.layers.Dense(128, activation='relu', input_shape=(100,)),
    tf.keras.layers.Dense(64, activation='relu'),
    tf.keras.layers.Dense(10, activation='softmax')
])

# Configuração do processo de treinamento
rede_neural.compile(optimizer='rmsprop',
                    loss='categorical_crossentropy',
                    metrics=['precision'])

# Preservação do modelo treinado
rede_neural.save('modelo_treinado.h5')

# Recuperação do modelo para uso posterior
modelo_recuperado = tf.keras.models.load_model('modelo_treinado.h5')

2. Salvamento e Carregamento com tf.train

A abordagem nativa do módulo tf.train oferece controle mais detalhado, sendo ideal para arquiteturas personalizadas.

import tensorflow as tf

# Definição de um grafo computacional dedicado
grafo_treinamento = tf.Graph()
with grafo_treinamento.as_default():
    # Placeholder para dados de entrada
    entrada_dados = tf.placeholder(tf.float32, shape=[None, 784])
    
    # Definição das camadas da rede
    camada_oculta = tf.layers.dense(entrada_dados, 256, activation=tf.nn.relu)
    saida_predicao = tf.layers.dense(camada_oculta, 10, activation=tf.nn.softmax)
    
    # Operador para salvamento de variáveis
   gerenciador_salvar = tf.train.Saver()

# Processo de treinamento e persistência
with tf.Session(graph=grafo_treinamento) as sess:
    sess.run(tf.global_variables_initializer())
    # Execução do treinamento
    gerenciador_salvar.save(sess, 'pesos_modelo.ckpt')

# Recuperação do modelo em um novo grafo
with tf.Graph().as_default() as grafo_inferencia:
    # Reconstrução da estrutura
    entrada = tf.placeholder(tf.float32, shape=[None, 784])
    hidden = tf.layers.dense(entrada, 256, activation=tf.nn.relu)
    predictions = tf.layers.dense(hidden, 10, activation=tf.nn.softmax)
    
    with tf.Session(graph=grafo_inferencia) as sess:
        restaurador = tf.train.import_meta_graph('pesos_modelo.ckpt.meta')
        restaurador.restore(sess, tf.train.latest_checkpoint('./'))

3. Usabilidade e Curva de Aprendizado

A interface tf.keras oferece maior produtividade para a maioria dos desenvolvedores. Com apenas duas linhas de código é possível preservar toda a estrutura do modelo, incluindo camadas, otimizador e estado de treinamento. O método model.save() gera um diretório completo contendo configuração, pesos e metadados do treinaemnto.

A abordagem tf.train exige maior familiaridade com o ecossistema TensorFlow. O desenvolvedor deve instanciar explicitamente um objeto Saver, definir quais variáveis serão preservadas e gerenciar separadamente a gravação do grafo computacioanl e dos parâmetros treinados.

4. Nível de Flexibilidade

O método tf.keras automatiza completamente o processo de serialização, o que representa uma vantagem para projetos com requisitos convencionais. Entretanto, essa abstração pode limitar cenários que demandam customização profunda, como salvamento parcial de camadas ou modificação dinâmica de grafos durante a restauração.

O método tf.train proporciona controle granular sobre cada aspecto do processo. É possível selecionar variáveis específicas para preservação, implementar lógicas customizadas durante a restauração e manipular diretamente o grafo computacional para adoptá-lo a diferentes contextos de execução.

Tags: tensorflow machine-learning deep-learning neural-networks keras

Publicado em 9-5 16:51