Метод set_weights класса Layer
Метод set_weights класса Layer
применяется к экземпляру слоя и позволяет
установить новые значения для всех весов слоя
одним вызовом. Первым параметром метод принимает
список массивов или тензоров, количество и форма
которых должны совпадать с количеством и формой
весов слоя. Вторым параметром можно передать
логическое значение, которое разрешает или
запрещает повторное создание переменных весов.
Метод возвращает None и изменяет веса
слоя на месте.
Синтаксис
layer.set_weights(weights)
layer.set_weights(weights, trainable)
Пример
Давайте создадим полносвязный слой Dense
с двумя нейронами и установим для него новые
веса при помощи метода set_weights:
import tensorflow as tf
import numpy as np
layer = tf.keras.layers.Dense(2, input_shape=(3,))
layer.build((None, 3))
new_weights = [
np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]),
np.array([0.5, 0.5])
]
layer.set_weights(new_weights)
print(layer.get_weights())
Результат выполнения кода:
[array([[1., 2.],
[3., 4.],
[5., 6.]], dtype=float32), array([0.5, 0.5], dtype=float32)]
Пример
Давайте создадим слой Dense, получим его
текущие веса методом get_weights, обнулим
их и вернём обратно через set_weights:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
layer = tf.keras.layers.Dense(2, input_shape=(3,))
layer.build((None, 3))
old_weights = layer.get_weights()
zero_weights = [np.zeros_like(w) for w in old_weights]
layer.set_weights(zero_weights)
print(layer.get_weights())
layer.set_weights(old_weights)
print(layer.get_weights()[1])
Результат выполнения кода:
[array([[0., 0.],
[0., 0.],
[0., 0.]], dtype=float32), array([0., 0.], dtype=float32)]
[0. 0.]
Пример
Давайте проверим, что метод set_weights
выбрасывает ошибку, если передать неподходящее
количество массивов весов:
import tensorflow as tf
import numpy as np
layer = tf.keras.layers.Dense(2, input_shape=(3,))
layer.build((None, 3))
try:
layer.set_weights([np.zeros((3, 2))])
except ValueError as e:
print("ValueError:", e)
Результат выполнения кода:
"ValueError: You called `set_weights(weights)` on layer \"dense\" with a weight list of length 1, but the layer was expecting 2 weights. Provided weights: [array([[0., 0.], ..."
Смотрите также
-
метод
get_weights,
который возвращает текущие веса слоя -
атрибут
weights,
который хранит список всех весов слоя -
метод
add_weight,
который добавляет новую переменную веса в слой -
метод
build,
который создает веса слоя на основе входной формы