TensorFlow - 模型导出
TensorFlow - 导出模型
Section titled “TensorFlow - 导出模型”导出训练好的 TensorFlow 模型是在各种环境中部署它的关键步骤,例如服务器、移动设备或 Web 浏览器。TensorFlow 的标准部署格式是 SavedModel。一个 SavedModel 包含了完整的 TensorFlow 程序,包括计算图、已学习的权重(变量)以及模型所需的任何资产(assets)。
对于 Keras 模型(在 TensorFlow 中常用),model.save() 方法是将模型导出为 SavedModel 格式的主要方式。你也可以保存为旧的 Keras H5 格式,但 SavedModel 通常更受青睐,因为它在 TensorFlow 生态系统内具有更广泛的兼容性。
以下是如何保存 Keras 模型然后加载回来的基本示例:
import tensorflow as tfimport 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 formatsaved_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 directoryloaded_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