Метод export
Метод export класса Model сохраняет модель в формате SavedModel.
Этот формат является универсальным для TensorFlow и позволяет загружать модель
в других средах выполнения, отличных от Python. Первым параметром метод принимает
путь к директории, в которую будет сохранена модель. Вторым необязательным
параметром можно передать verbose для управления выводом информации
в процессе экспорта.
Экспортированная модель содержит граф вычислений и веса, что позволяет использовать ее без исходного Python-кода. Это особенно удобно при развертывании моделей в production-среде.
Синтаксис
model.export(filepath, [verbose])
Пример
Давайте создадим простую модель и экспортируем ее в директорию
'my_model':
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(10, activation='relu', input_shape=(5,)),
tf.keras.layers.Dense(1)
])
model.compile(optimizer='adam', loss='mse')
model.export('my_model')
print("Model exported successfully")
Результат выполнения кода:
"Model exported successfully"
Пример
Давайте обучим модель на простых данных и экспортируем ее
с параметром verbose=True:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(10, activation='relu', input_shape=(5,)),
tf.keras.layers.Dense(1)
])
model.compile(optimizer='adam', loss='mse')
x = tf.constant([[1.0, 2.0, 3.0, 4.0, 5.0]])
y = tf.constant([[10.0]])
model.fit(x, y, epochs=1, verbose=0)
model.export('trained_model', verbose=True)
Результат выполнения кода:
"Saved artifact at 'trained_model'. The following endpoints are available: ..."
Пример
Давайте загрузим экспортированную модель обратно с помощью
tf.saved_model.load и выполним предсказание:
import tensorflow as tf
loaded = tf.saved_model.load('trained_model')
x = tf.constant([[1.0, 2.0, 3.0, 4.0, 5.0]])
res = loaded(x)
print(res)
Результат выполнения кода:
tf.Tensor([[10.0]], shape=(1, 1), dtype=float32)
Смотрите также
-
метод
save,
который сохраняет модель в формате Keras -
метод
save_weights,
который сохраняет только веса модели -
метод
load_weights,
который загружает веса в модель -
метод
to_json,
который сериализует архитектуру модели в JSON