Метод set_weights
Метод set_weights класса Model применяется к модели
для установки ее весовых коэффициентов. Первым параметром
метод принимает список массивов NumPy, значения которых
будут присвоены весам модели в порядке их следования.
Количество массивов и их формы должны соответствовать
структуре весов модели, полученных через метод get_weights.
Метод не возвращает значение. Он изменяет веса модели на месте. Если переданные массивы не соответствуют ожидаемым формам или количеству, будет вызвано исключение.
Синтаксис
model.set_weights(weights)
Пример
Давайте создадим простую модель, получим ее веса,
обнулим их и установим обратно через метод set_weights:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(2,))
])
original_weights = model.get_weights()
print("Original weights:")
for w in original_weights:
print(w)
zero_weights = [np.zeros_like(w) for w in original_weights]
model.set_weights(zero_weights)
print("After set_weights with zeros:")
for w in model.get_weights():
print(w)
Результат выполнения кода:
Original weights:
[[ 0.86039335 -0.21362835]
[ 0.45075434 -0.42822015]]
[0. 0.]
After set_weights with zeros:
[[0. 0.]
[0. 0.]]
[0. 0.]
Пример
Давайте восстановим исходные веса модели после их изменения:
<+python+>
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
saved_weights = model.get_weights()
new_weights = [np.ones_like(w) for w in saved_weights]
model.set_weights(new_weights)
print("After set_weights with ones:")
for w in model.get_weights():
print(w)
model.set_weights(saved_weights)
print("After restoring original weights:")
for w in model.get_weights():
print(w)
<-python+>
Результат выполнения кода:
<+python+>
After set_weights with ones:
[[1. 1. 1.]
[1. 1. 1.]]
[1. 1. 1.]
After restoring original weights:
[[-0.02856663 0.04454689 0.0236021 ]
[ 0.06184066 -0.03228709 0.02872896]]
[0. 0. 0.]
<-python+>
Смотрите также
-
метод
get_weights,
который возвращает список весов модели -
метод
save_weights,
который сохраняет веса модели в файл -
метод
load_weights,
который загружает веса модели из файла -
атрибут
weights,
который содержит список всех весов модели