Метод save_weights
Метод save_weights применяется к экземпляру
модели Model и сохраняет значения
всех весов (обучаемых и необучаемых) в указанный
файл или директорию. Первым параметром метод
принимает путь к файлу или папке. Вторым
необязательным параметром можно передать
формат сохранения, например 'tf' для
формата TensorFlow или 'h5' для формата
HDF5. Также можно передать аргумент
overwrite, который разрешает или
запрещает перезапись существующих файлов.
Синтаксис
model.save_weights(filepath, [overwrite], [save_format])
Пример
Давайте создадим простую модель, обучим её на небольших данных и сохраним веса в файл формата TensorFlow:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,)),
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
x = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
y = tf.constant([[1.0], [2.0], [3.0]])
model.fit(x, y, epochs=1, verbose=0)
model.save_weights('model.weights.h5')
print("Weights saved")
Результат выполнения кода:
"Weights saved"
Пример
Давайте сохраним веса модели в формате HDF5,
явно указав параметр save_format:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,)),
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
x = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
y = tf.constant([[1.0], [2.0], [3.0]])
model.fit(x, y, epochs=1, verbose=0)
model.save_weights('model.weights.h5', save_format='h5')
print("Weights saved in HDF5 format")
Результат выполнения кода:
"Weights saved in HDF5 format"
Пример
Давайте сохраним веса в директорию, запретив перезапись существующих файлов:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,)),
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
x = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
y = tf.constant([[1.0], [2.0], [3.0]])
model.fit(x, y, epochs=1, verbose=0)
model.save_weights('weights_dir', overwrite=False)
print("Weights saved to directory")
Результат выполнения кода:
"Weights saved to directory"
Смотрите также
-
метод
save,
который сохраняет всю модель целиком -
метод
load_weights,
который загружает веса модели из файла -
метод
get_weights,
который возвращает веса модели в виде списка -
метод
set_weights,
который устанавливает веса модели из списка