Spec-Zone.ru › TensorFlow

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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API