Атрибут trainable_weights класса Model
Атрибут trainable_weights класса Model возвращает
список тензоров, которые являются обучаемыми параметрами
модели. К ним относятся веса и смещения слоев, у которых
свойство trainable установлено в True. Именно
эти тензоры обновляются оптимизатором в процессе вызова
методов fit и train_on_batch. Атрибут не
принимает параметров и доступен только для чтения.
Атрибут полезен, когда нужно вручную изучить обучаемые параметры, передать их во внешний оптимизатор или проверить, какие слои действительно участвуют в обучении.
Синтаксис
model.trainable_weights
Пример
Давайте создадим простую модель с одним полносвязным слоем
и посмотрим, какие тензоры попадают в атрибут
trainable_weights:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
print(model.trainable_weights)
Результат выполнения кода:
[<tf.Variable 'dense/kernel:0' shape=(2, 3) dtype=float32, numpy=
array([[ 0.6846645 , 0.7841382 , -0.22484612],
[-0.8302233 , -0.17961967, 0.13435602]], dtype=float32)>, <tf.Variable 'dense/bias:0' shape=(3,) dtype=float32, numpy=array([0., 0., 0.], dtype=float32)>]
В список попали два тензора: ядро слоя kernel и
смещение bias.
Пример
Давайте посмотрим, как меняется содержимое атрибута, если
заморозить слой, установив свойство trainable в
False:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
model.layers[0].trainable = False
print(model.trainable_weights)
print(model.non_trainable_weights)
Результат выполнения кода:
[]
[<tf.Variable 'dense/kernel:0' shape=(2, 3) dtype=float32, numpy=
array([[ 0.6846645 , 0.7841382 , -0.22484612],
[-0.8302233 , -0.17961967, 0.13435602]], dtype=float32)>, <tf.Variable 'dense/bias:0' shape=(3,) dtype=float32, numpy=array([0., 0., 0.], dtype=float32)>]
После заморозки слоя атрибут trainable_weights стал
пустым, а все параметры перешли в
non_trainable_weights.
Пример
Давайте убедимся, что атрибут trainable_weights
действительно содержит те же объекты, что и метод
get_weights для обучаемых переменных:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
res = len(model.trainable_weights)
print(res)
for w in model.trainable_weights:
print(w.name, w.shape)
Результат выполнения кода:
2
dense/kernel:0 (2, 3)
dense/bias:0 (3,)
Смотрите также
-
атрибут
weights,
который возвращает все веса модели -
атрибут
non_trainable_weights,
который возвращает необучаемые веса модели -
метод
get_weights,
который возвращает значения весов модели -
метод
set_weights,
который устанавливает значения весов модели