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

Атрибут trainable

Атрибут trainable класса Model управляет тем, будут ли переменные модели (веса и смещения) обновляться в процессе обучения. Если атрибут установлен в True, то слои модели доступны для тренировки, и их веса будут изменяться при вызове методов fit или train_on_batch. Если установить значение False, то все веса модели замораживаются: они перестают обновляться, но по-прежнему используются для прямого прохода. Это полезно при тонкой настройке (fine-tuning), когда нужно обучить только верхние слои, не меняя предварительно обученные веса.

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

Синтаксис

model.trainable = True # или False

Пример

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

import tensorflow as tf tf.random.set_seed(0) # Create a simple model model = tf.keras.Sequential([ tf.keras.layers.Dense(2, input_shape=(2,)), tf.keras.layers.Dense(1) ]) # Compile and train model.compile(optimizer='sgd', loss='mse') model.fit([[1, 2], [3, 4]], [[5], [6]], epochs=1, verbose=0) # Save weights before freezing weights_before = model.get_weights() # Freeze the model model.trainable = False # Try to train again model.fit([[1, 2], [3, 4]], [[5], [6]], epochs=1, verbose=0) # Check weights after freezing weights_after = model.get_weights() # Compare print("Weights changed:", not all( tf.reduce_all(tf.equal(w1, w2)) for w1, w2 in zip(weights_before, weights_after) ))

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

Weights changed: False

Пример

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

import tensorflow as tf model = tf.keras.Sequential([ tf.keras.layers.Dense(4, input_shape=(2,)), tf.keras.layers.Dense(1) ]) # By default all layers are trainable print("Layer 1 trainable:", model.layers[0].trainable) print("Layer 2 trainable:", model.layers[1].trainable) # Freeze only the first layer model.layers[0].trainable = False print("After freezing layer 1:") print("Layer 1 trainable:", model.layers[0].trainable) print("Layer 2 trainable:", model.layers[1].trainable)

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

Layer 1 trainable: True Layer 2 trainable: True After freezing layer 1: Layer 1 trainable: False Layer 2 trainable: True

Пример

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

import tensorflow as tf model = tf.keras.Sequential([ tf.keras.layers.Dense(4, input_shape=(2,)), tf.keras.layers.Dense(1) ]) print("Trainable weights count:", len(model.trainable_weights)) print("Non-trainable weights count:", len(model.non_trainable_weights)) # Freeze the whole model model.trainable = False print("After freezing:") print("Trainable weights count:", len(model.trainable_weights)) print("Non-trainable weights count:", len(model.non_trainable_weights))

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

Trainable weights count: 4 Non-trainable weights count: 0 After freezing: Trainable weights count: 0 Non-trainable weights count: 4

Смотрите также

  • класс Model,
    который представляет собой базовый класс для моделей
  • метод fit,
    который обучает модель на данных
  • атрибут trainable_weights,
    который возвращает список обучаемых весов
  • атрибут non_trainable_weights,
    который возвращает список необучаемых весов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить