РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
640 of 824 menu

Функция saved_model.load

Функция saved_model.load загружает модель, ранее сохраненную с помощью функции saved_model.save. Первым параметром функция принимает путь к директории, в которой находится сохраненная модель. Вторым необязательным параметром можно передать объект LoadOptions, который управляет процессом загрузки, например, позволяет указать конкретные объекты для загрузки. Функция возвращает объект AutoTrackable, содержащий загруженные функции и переменные модели.

Синтаксис

tf.saved_model.load(export_dir, [options])

Пример

Давайте создадим простую модель, сохраним ее, а затем загрузим обратно:

import tensorflow as tf class MyModel(tf.Module): def __init__(self): super(MyModel, self).__init__() self.v = tf.Variable(3.0) @tf.function(input_signature=[tf.TensorSpec(shape=None, dtype=tf.float32)]) def __call__(self, x): return x * self.v model = MyModel() tf.saved_model.save(model, 'my_model') loaded = tf.saved_model.load('my_model') res = loaded(tf.constant(5.0)) print(res)

Результат выполнения кода:

tf.Tensor(15.0, shape=(), dtype=float32)

Пример

Давайте загрузим модель и получим доступ к ее переменным и сигнатурам:

import tensorflow as tf class MyModel(tf.Module): def __init__(self): super(MyModel, self).__init__() self.v = tf.Variable(3.0) @tf.function(input_signature=[tf.TensorSpec(shape=None, dtype=tf.float32)]) def __call__(self, x): return x * self.v model = MyModel() tf.saved_model.save(model, 'my_model') loaded = tf.saved_model.load('my_model') print(list(loaded.signatures.keys())) print(loaded.v.numpy())

Результат выполнения кода:

['__saved_model_init_op', '__call__'] 3.0

Пример

Давайте загрузим модель с помощью объекта LoadOptions, указав конкретную сигнатуру:

<+python+> import tensorflow as tf class MyModel(tf.Module): def __init__(self): super(MyModel, self).__init__() self.v = tf.Variable(3.0) @tf.function(input_signature=[tf.TensorSpec(shape=None, dtype=tf.float32)]) def __call__(self, x): return x * self.v model = MyModel() tf.saved_model.save(model, 'my_model') options = tf.saved_model.LoadOptions(experimental_io_device='/job:localhost') loaded = tf.saved_model.load('my_model', options=options) res = loaded(tf.constant(5.0)) print(res) <-python+>

Результат выполнения кода:

tf.Tensor(15.0, shape=(), dtype=float32)

Смотрите также

  • функцию saved_model.save,
    которая сохраняет модель в указанную директорию
  • класс LoadOptions,
    который управляет параметрами загрузки модели
  • класс SaveOptions,
    который управляет параметрами сохранения модели
  • функцию latest_checkpoint,
    которая находит последний сохраненный чекпоинт
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить