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

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