Атрибут trainable
Атрибут trainable класса Layer
управляет тем, будут ли переменные слоя
обновляться во время обучения. Если атрибут
установлен в True, веса слоя участвуют
в обучении. Если в False, веса
замораживаются и не изменяются. Атрибут
можно задать при создании слоя или изменить
после его создания.
Синтаксис
layer = tf.keras.layers.Dense(units, trainable=True)
layer.trainable = False
Пример
Давайте создадим слой Dense с атрибутом
trainable равным True и проверим
значение атрибута:
import tensorflow as tf
layer = tf.keras.layers.Dense(3, trainable=True)
print(layer.trainable)
Результат выполнения кода:
True
Пример
Давайте создадим слой с атрибутом
trainable равным False:
import tensorflow as tf
layer = tf.keras.layers.Dense(3, trainable=False)
print(layer.trainable)
Результат выполнения кода:
False
Пример
Давайте создадим слой, а затем изменим
атрибут trainable после его создания:
import tensorflow as tf
layer = tf.keras.layers.Dense(3)
print(layer.trainable)
layer.trainable = False
print(layer.trainable)
Результат выполнения кода:
True
False
Пример
Давайте создадим модель с двумя слоями, заморозим первый слой и посмотрим на список обучаемых переменных:
Результат выполнения кода:
False
True
Пример
Давайте посмотрим, как атрибут trainable
влияет на список обучаемых весов слоя:
import tensorflow as tf
layer = tf.keras.layers.Dense(3, input_shape=(4,))
layer.build((None, 4))
print(len(layer.trainable_weights))
layer.trainable = False
print(len(layer.trainable_weights))
print(len(layer.non_trainable_weights))
Результат выполнения кода:
2
0
2
Смотрите также
-
класс
Layer,
который является базовым классом для всех слоев -
атрибут
trainable_weights,
который содержит список обучаемых весов слоя -
атрибут
non_trainable_weights,
который содержит список необучаемых весов слоя -
атрибут
weights,
который содержит все веса слоя