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

Атрибут 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

Пример

Давайте создадим модель с двумя слоями, заморозим первый слой и посмотрим на список обучаемых переменных:

<+python+> import tensorflow as tf model = tf.keras.Sequential([ tf.keras.layers.Dense(4, input_shape=(3,)), tf.keras.layers.Dense(2) ]) model.layers[0].trainable = False for layer in model.layers: print(layer.trainable) <-python+>

Результат выполнения кода:

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