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

Метод 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,
    который создает веса слоя на основе входной формы
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить