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.