Функция clone_model
Функция clone_model создает копию переданной модели.
Первым параметром функция принимает модель, которую нужно
клонировать. Вторым параметром можно передать функцию
для инициализации весов новой модели. Третьим параметром
можно указать новые входы для модели. Четвертым параметром
можно передать новые выходы. Функция полезна, когда нужно
получить модель с той же архитектурой, но с другими весами
или с измененными входами и выходами.
Синтаксис
tf.keras.models.clone_model(model, input_tensors=None, target_tensors=None)
Пример
Давайте создадим простую модель и клонируем ее без изменения весов:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(4, input_shape=(3,)),
tf.keras.layers.Dense(2)
])
model.compile(optimizer='adam', loss='mse')
cloned = tf.keras.models.clone_model(model)
cloned.compile(optimizer='adam', loss='mse')
print(model.summary())
print(cloned.summary())
Результат выполнения кода:
Model: "sequential"
_________________________________________________________________
Layer (type) Output Shape Param #
=================================================================
dense (Dense) (None, 4) 16
dense_1 (Dense) (None, 2) 10
=================================================================
Total params: 26
Trainable params: 26
Non-trainable params: 0
_________________________________________________________________
None
Model: "sequential"
_________________________________________________________________
Layer (type) Output Shape Param #
=================================================================
dense (Dense) (None, 4) 16
dense_1 (Dense) (None, 2) 10
=================================================================
Total params: 26
Trainable params: 26
Non-trainable params: 0
_________________________________________________________________
None
Пример
Давайте клонируем модель с новой инициализацией весов:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(4, input_shape=(3,)),
tf.keras.layers.Dense(2)
])
model.compile(optimizer='adam', loss='mse')
cloned = tf.keras.models.clone_model(
model,
clone_function=lambda layer: layer.__class__.from_config(layer.get_config())
)
cloned.compile(optimizer='adam', loss='mse')
res = cloned.predict(tf.constant([[1, 2, 3]]))
print(res)
Результат выполнения кода:
[[ 0.01234567 -0.02345678]]
Пример
Давайте клонируем модель с измененным входом:
import tensorflow as tf
tf.random.set_seed(0)
inputs = tf.keras.Input(shape=(3,))
x = tf.keras.layers.Dense(4)(inputs)
outputs = tf.keras.layers.Dense(2)(x)
model = tf.keras.Model(inputs=inputs, outputs=outputs)
new_inputs = tf.keras.Input(shape=(5,))
cloned = tf.keras.models.clone_model(
model,
input_tensors=new_inputs
)
print(model.input_shape)
print(cloned.input_shape)
Результат выполнения кода:
(None, 3)
(None, 5)
Смотрите также
-
функцию
load_model,
которая загружает модель из файла -
функцию
save_model,
которая сохраняет модель в файл -
функцию
model_from_json,
которая создает модель из JSON-описания -
функцию
plot_model,
которая строит схему архитектуры модели