Метод load_weights
Метод load_weights класса Model загружает веса модели из файла, созданного методом save_weights, или из списка массивов NumPy. Первым параметром метод принимает путь к файлу с весами или список массивов. Вторым параметром можно передать флаг by_name, который указывает, загружать ли веса по именам слоев, а не по порядку. Третьим параметром можно передать флаг skip_mismatch, который позволяет пропускать слои с несовпадающей формой весов. Метод возвращает объект модели, что позволяет строить цепочки вызовов.
Синтаксис
model.load_weights(filepath, [by_name], [skip_mismatch])
Пример
Давайте создадим простую модель, сохраним ее веса в файл, а затем загрузим их обратно:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
model.save_weights('model.weights.h5')
new_model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
new_model.load_weights('model.weights.h5')
res = new_model.get_weights()
print(res)
Результат выполнения кода:
[array([[ 0.51020974, -0.2746159 ],
[ 0.24904633, 0.5070965 ],
[-0.31243873, 0.39328647]], dtype=float32), array([0., 0.], dtype=float32)]
Пример
Давайте загрузим веса по именам слоев, используя параметр by_name:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,), name='dense_first')
])
model.save_weights('model.weights.h5')
new_model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,), name='dense_first')
])
new_model.load_weights('model.weights.h5', by_name=True)
res = new_model.get_weights()
print(res)
Результат выполнения кода:
[array([[ 0.51020974, -0.2746159 ],
[ 0.24904633, 0.5070965 ],
[-0.31243873, 0.39328647]], dtype=float32), array([0., 0.], dtype=float32)]
Пример
Давайте загрузим веса из списка массивов NumPy:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
weights = [
np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]),
np.array([0.5, 0.5])
]
model.load_weights(weights)
res = model.get_weights()
print(res)
Результат выполнения кода:
[array([[1., 2.],
[3., 4.],
[5., 6.]], dtype=float32), array([0.5, 0.5], dtype=float32)]
Смотрите также
-
метод
save_weights,
который сохраняет веса модели в файл -
метод
get_weights,
который возвращает веса модели в виде списка массивов -
метод
set_weights,
который устанавливает веса модели из списка массивов -
метод
save,
который сохраняет модель целиком в файл