Skip to content

TensorFlow - 模型导出

导出训练好的 TensorFlow 模型是在各种环境中部署它的关键步骤,例如服务器、移动设备或 Web 浏览器。TensorFlow 的标准部署格式是 SavedModel。一个 SavedModel 包含了完整的 TensorFlow 程序,包括计算图、已学习的权重(变量)以及模型所需的任何资产(assets)。

对于 Keras 模型(在 TensorFlow 中常用),model.save() 方法是将模型导出为 SavedModel 格式的主要方式。你也可以保存为旧的 Keras H5 格式,但 SavedModel 通常更受青睐,因为它在 TensorFlow 生态系统内具有更广泛的兼容性。

以下是如何保存 Keras 模型然后加载回来的基本示例:

import tensorflow as tf
import numpy as np
# 1. Define and compile a simple Keras model (example)
model = tf.keras.Sequential([
tf.keras.layers.Dense(128, activation='relu', input_shape=(784,)),
tf.keras.layers.Dropout(0.2),
tf.keras.layers.Dense(10, activation='softmax')
])
model.compile(optimizer='adam',
loss='sparse_categorical_crossentropy',
metrics=['accuracy'])
# (Imagine the model is trained here with some data)
# Example dummy training:
# x_dummy = np.random.rand(100, 784)
# y_dummy = np.random.randint(0, 10, 100)
# model.fit(x_dummy, y_dummy, epochs=1)
# 2. Save the model to a directory in SavedModel format
saved_model_path = './my_simple_model'
model.save(saved_model_path)
print(f"Model saved to {saved_model_path}")
# 3. Load the model from the SavedModel directory
loaded_model = tf.keras.models.load_model(saved_model_path)
print("Model loaded successfully.")
# Verify the loaded model (optional)
# loaded_model.summary()
# You can now use loaded_model for inference, e.g., loaded_model.predict(...)
# For non-Keras models or custom TensorFlow modules/functions,
# you can use `tf.saved_model.save()` and `tf.saved_model.load()`.
# Example:
# tf.saved_model.save(your_tf_module, './my_custom_tf_module_saved_model')
# loaded_custom_module = tf.saved_model.load('./my_custom_tf_module_saved_model')

一旦模型是 SavedModel 格式,它就可以被:

  • 使用 TensorFlow Serving 进行服务,用于生产环境。
  • 转换为 TensorFlow Lite(.tflite)格式,用于部署到移动和嵌入式设备。这通常涉及优化,例如量化,以减小模型大小并提高推理速度。
  • 与 TensorFlow.js 一起使用,直接在 Web 浏览器或 Node.js 应用程序中运行模型。
  • 部署在支持 TensorFlow 模型的各种云平台和边缘设备上。

有关模型保存和序列化的更多信息,请参阅 TensorFlow 官方文档:https://www.tensorflow.org/guide/saved_model