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

Метод save класса Checkpoint

Метод save класса Checkpoint сохраняет текущие значения отслеживаемых объектов в контрольную точку. Первым параметром метод принимает путь к файлу или каталогу, куда будет сохранено состояние. Вторым параметром можно передать сессию для сохранения. Метод возвращает путь, по которому была сохранена контрольная точка. Класс Checkpoint позволяет сохранять и восстанавливать состояние моделей, оптимизаторов и других объектов TensorFlow.

Синтаксис

tf.train.Checkpoint.save(file_prefix, [session])

Пример

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

import tensorflow as tf model = tf.keras.Sequential([ tf.keras.layers.Dense(5, input_shape=(3,)) ]) optimizer = tf.keras.optimizers.Adam() checkpoint = tf.train.Checkpoint( model=model, optimizer=optimizer ) save_path = checkpoint.save('./ckpt/test') print(save_path)

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

"./ckpt/test-1"

Пример

Давайте сохраним состояние модели после обучения на данных:

import tensorflow as tf tf.random.set_seed(0) model = tf.keras.Sequential([ tf.keras.layers.Dense(2, input_shape=(3,)) ]) optimizer = tf.keras.optimizers.SGD(learning_rate=0.01) checkpoint = tf.train.Checkpoint( model=model, optimizer=optimizer ) x = tf.constant([[1, 2, 3], [4, 5, 6]]) y = tf.constant([[1, 0], [0, 1]]) with tf.GradientTape() as tape: predictions = model(x) loss = tf.keras.losses.mean_squared_error(y, predictions) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) save_path = checkpoint.save('./ckpt/trained') print(save_path)

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

"./ckpt/trained-1"

Пример

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

import tensorflow as tf model = tf.keras.Sequential([ tf.keras.layers.Dense(3, input_shape=(2,)) ]) optimizer = tf.keras.optimizers.Adam() checkpoint = tf.train.Checkpoint( model=model, optimizer=optimizer ) path1 = checkpoint.save('./ckpt/step1') path2 = checkpoint.save('./ckpt/step2') print(path1) print(path2)

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

"./ckpt/step1-1" "./ckpt/step2-1"

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

  • класс Checkpoint,
    который управляет сохранением и восстановлением состояния
  • метод restore,
    который восстанавливает состояние из контрольной точки
  • метод read,
    который читает состояние из контрольной точки без восстановления
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить