Атрибут 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,
который возвращает список необучаемых весов