tf.keras.export.ExportArchive
ExportArchive используется для записи артефактов SavedModel (например, для инференса).
tf.keras.export.ExportArchive()
Если у вас есть модель или слой Keras, которые вы хотите экспортировать в SavedModel для предоставления сервиса (например, через TensorFlow-Serving), вы можете использовать ExportArchive для настройки различных точек доступа к сервису, которые вам нужно предоставить, а также их сигнатур. Просто создайте экземпляр ExportArchive, используйте track() для регистрации слоев или моделей, которые нужно использовать, а затем используйте метод add_endpoint() для регистрации новой точки доступа к сервису. По завершении используйте метод write_out() для сохранения артефакта.
Полученный артефакт — это SavedModel, который можно перезагрузить с помощью tf.saved_model.load.
Примеры:
Вот как экспортировать модель для инференса.
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(
name="serve",
fn=model.call,
input_signature=[tf.TensorSpec(shape=(None, 3), dtype=tf.float32)],
)
export_archive.write_out("path/to/location")
# Elsewhere, we can reload the artifact and serve it.
# The endpoint we added is available as a method:
serving_model = tf.saved_model.load("path/to/location")
outputs = serving_model.serve(inputs)
Вот как экспортировать модель с одним выходом для инференса и одним выходом для прохода в режиме обучения (например, с дропаутом).
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(
name="call_inference",
fn=lambda x: model.call(x, training=False),
input_signature=[tf.TensorSpec(shape=(None, 3), dtype=tf.float32)],
)
export_archive.add_endpoint(
name="call_training",
fn=lambda x: model.call(x, training=True),
input_signature=[tf.TensorSpec(shape=(None, 3), dtype=tf.float32)],
)
export_archive.write_out("path/to/location")
Примечание о отслеживании ресурсов:
ExportArchive может автоматически отслеживать все tf.Variables, используемые его точками доступа, поэтому в большинстве случаев вызов .track(model) не строго обязателен. Однако, если ваша модель использует слои поиска, такие как IntegerLookup, StringLookup или TextVectorization, необходимо явно отслеживать их с помощью .track(model).
Явное отслеживание также необходимо, если вам нужно получить доступ к свойствам variables, trainable_variables или non_trainable_variables в восстановленном архиве.
| Атрибуты | |
|---|---|
non_trainable_variables | |
trainable_variables | |
variables | |
Методы
add_endpoint
add_endpoint(
name, fn, input_signature=None, jax2tf_kwargs=None
)
Регистрация новой точки доступа к сервису.
| Аргументы | |
|---|---|
name | Строка, имя точки доступа. |
fn | Функция. Она должна использовать только ресурсы (например, объекты tf.Variable или объекты tf.lookup.StaticHashTable), доступные в отслеженных моделях/слоях через ExportArchive (вы можете вызвать .track(model) для отслеживания новой модели). Формат и тип данных входных данных для функции должны быть известны. Для этого вы можете либо 1) убедиться, что fn — это tf.function, который был вызван хотя бы один раз, или 2) предоставить аргумент input_signature, который указывает формат и тип данных входных данных (см. ниже). |
input_signature | Используется для указания формата и типа данных входных данных для fn. Список объектов tf.TensorSpec (по одному на позиционный аргумент входа fn). Вложенные аргументы разрешены (см. ниже пример функциональной модели с 2 аргументами ввода). |
jax2tf_kwargs | Необязательно. Словарь для аргументов, передаваемых в jax2tf. Поддерживается только в случае использования JAX в качестве бэкенда. Смотрите документацию по jax2tf.convert. Значения для native_serialization и polymorphic_shapes, если они не указаны, вычисляются автоматически. |
| Возвращает | |
|---|---|
tf.function — обертка fn, добавленная в архив. |
Пример:
Добавление точки доступа, используя аргумент input_signature, когда модель имеет один аргумент ввода:
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(
name="serve",
fn=model.call,
input_signature=[tf.TensorSpec(shape=(None, 3), dtype=tf.float32)],
)
Добавление точки доступа, используя аргумент input_signature, когда модель имеет два позиционных аргумента ввода:
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(
name="serve",
fn=model.call,
input_signature=[
tf.TensorSpec(shape=(None, 3), dtype=tf.float32),
tf.TensorSpec(shape=(None, 4), dtype=tf.float32),
],
)
Добавление точки доступа, используя аргумент input_signature, когда модель имеет один аргумент ввода, который представляет собой список из 2 тензоров (например, функциональная модель с 2 входами):
model = keras.Model(inputs=[x1, x2], outputs=outputs)
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(
name="serve",
fn=model.call,
input_signature=[
[
tf.TensorSpec(shape=(None, 3), dtype=tf.float32),
tf.TensorSpec(shape=(None, 4), dtype=tf.float32),
],
],
)
Это также работает с входными данными в виде словаря:
model = keras.Model(inputs={"x1": x1, "x2": x2}, outputs=outputs)
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(
name="serve",
fn=model.call,
input_signature=[
{
"x1": tf.TensorSpec(shape=(None, 3), dtype=tf.float32),
"x2": tf.TensorSpec(shape=(None, 4), dtype=tf.float32),
},
],
)
Добавление точки доступа, которая является tf.function:
@tf.function()
def serving_fn(x):
return model(x)
# The function must be traced, i.e. it must be called at least once.
serving_fn(tf.random.normal(shape=(2, 3)))
export_archive = ExportArchive()
export_archive.track(model)
export_archive.add_endpoint(name="serve", fn=serving_fn)
add_variable_collection
add_variable_collection(
name, variables
)
Регистрация набора переменных для извлечения после перезагрузки.
| Аргументы | |
|---|---|
name | Имя коллекции в виде строки. |
variables | Кортеж/список/множество экземпляров tf.Variable. |
Пример:
export_archive = ExportArchive()
export_archive.track(model)
# Register an endpoint
export_archive.add_endpoint(
name="serve",
fn=model.call,
input_signature=[tf.TensorSpec(shape=(None, 3), dtype=tf.float32)],
)
# Save a variable collection
export_archive.add_variable_collection(
name="optimizer_variables", variables=model.optimizer.variables)
export_archive.write_out("path/to/location")
# Reload the object
revived_object = tf.saved_model.load("path/to/location")
# Retrieve the variables
optimizer_variables = revived_object.optimizer_variables
track
track(
resource
)
Отслеживание переменных (и других активов) слоя или модели.
По умолчанию все переменные, используемые функцией точки доступа, автоматически отслеживаются при вызове add_endpoint(). Однако не переменные активы, такие как таблицы поиска, необходимо отслеживать вручную. Обратите внимание, что таблицы поиска, используемые встроенными слоями Keras (TextVectorization, IntegerLookup, StringLookup), автоматически отслеживаются в add_endpoint().
| Аргументы | |
|---|---|
resource | Отслеживаемый ресурс TensorFlow. |
write_out
write_out(
filepath, options=None
)
Запись соответствующего SavedModel на диск.
| Аргументы | |
|---|---|
filepath | str или pathlib.Path объект. Путь для сохранения артефакта. |
options | Объект tf.saved_model.SaveOptions, определяющий параметры сохранения SavedModel. |
Примечание для TF-Serving: все точки доступа, зарегистрированные через add_endpoint(), становятся видимыми для TF-Serving в артефакте SavedModel. Кроме того, первая зарегистрированная точка доступа становится видимой под псевдонимом "serving_default" (если вручную не была зарегистрирована точка доступа с именем "serving_default"), так как TF-Serving требует наличия этой точки доступа.
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/api_docs/python/tf/keras/export/ExportArchive