Атрибут non_trainable_weights
Атрибут non_trainable_weights класса Model
содержит список тензоров, которые не подлежат
обновлению во время обучения. К таким тензорам
относятся, например, статистики слоев нормализации
или другие переменные, которые модель использует,
но не оптимизирует. Атрибут не принимает параметров
и доступен только для чтения. Он возвращает список
объектов tf.Variable.
Синтаксис
model.non_trainable_weights
Пример
Давайте создадим простую модель с одним слоем
Dense и посмотрим на список
нетренируемых весов:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
res = model.non_trainable_weights
print(res)
Результат выполнения кода:
[]
Как видим, у обычного полносвязного слоя нет нетренируемых весов, поэтому список пуст.
Пример
Давайте создадим модель со слоем
BatchNormalization, который содержит
нетренируемые веса:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,)),
tf.keras.layers.BatchNormalization()
])
res = model.non_trainable_weights
print(len(res))
for w in res:
print(w.name, w.shape)
Результат выполнения кода:
2
"batch_normalization/moving_mean:0 (3,)"
"batch_normalization/moving_variance:0 (3,)"
Слой BatchNormalization добавил две
нетренируемые переменные: скользящее среднее и
скользящую дисперсию.
Пример
Давайте сравним общее количество весов модели с количеством тренируемых и нетренируемых весов:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,)),
tf.keras.layers.BatchNormalization()
])
print("All weights:", len(model.weights))
print("Trainable:", len(model.trainable_weights))
print("Non-trainable:", len(model.non_trainable_weights))
Результат выполнения кода:
"All weights: 6"
"Trainable: 4"
"Non-trainable: 2"
У модели оказалось 6 весов: 4 тренируемых (ядро и смещение полносвязного слоя, а также гамма и бета слоя нормализации) и 2 нетренируемых (скользящее среднее и скользящая дисперсия).
Пример
Давайте заморозим слой и посмотрим, как изменятся списки весов:
Результат выполнения кода:
"Trainable: 2"
"Non-trainable: 4"
После заморозки первого слоя его веса переместились из списка тренируемых в список нетренируемых.
Смотрите также
-
атрибут
trainable_weights,
который возвращает список тренируемых весов модели -
атрибут
weights,
который возвращает список всех весов модели -
метод
get_weights,
который возвращает значения весов модели -
метод
summary,
который выводит сводку по слоям и параметрам модели