You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow中能否将已训练模型的部分加载到新模型且免重训?

Can I reuse parts of a deprecated trained model in a new model with a different architecture in TensorFlow?

Absolutely! This is a standard transfer learning scenario, and TensorFlow makes it straightforward to pull off. You can load the useful weights from your old model into corresponding layers of your new model, freeze those reused layers so they don’t get retrained, and train the rest of the new model from scratch.

Here’s a step-by-step breakdown with practical code examples:

1. Load your old model/weights

First, get access to the weights of the layers you want to reuse. You can either load the entire old model, or rebuild its architecture and load just the weights if you don’t have the full model saved.

import tensorflow as tf

# Option 1: Load the full deprecated SavedModel
old_model = tf.keras.models.load_model('/path/to/your/old_model')

# Option 2: Rebuild old model architecture then load weights (if you only have .h5 weights)
# def build_old_model():
#     inputs = tf.keras.Input(shape=(28,28,1))
#     x = tf.keras.layers.Conv2D(32, (3,3), activation='relu')(inputs)
#     # ... replicate the rest of the old model's layers ...
#     outputs = tf.keras.layers.Dense(10, activation='softmax')(x)
#     return tf.keras.Model(inputs=inputs, outputs=outputs)
# old_model = build_old_model()
# old_model.load_weights('/path/to/old_model_weights.h5')

2. Build your new model with reused layers

There are two common approaches to reuse old layers: directly plugging them into the new model, or copying their weights into fresh layers with matching structures.

Option A: Directly reuse old model layers

If the layers you want to keep have exact input/output shapes that fit your new model, just plug them in and freeze them.

# Select the layers to reuse (e.g., first 3 layers of the old model)
reused_layers = old_model.layers[:3]

# Freeze these layers to prevent weight updates during training
for layer in reused_layers:
    layer.trainable = False

# Construct the new model
inputs = tf.keras.Input(shape=(your_new_input_shape,))
x = inputs

# Add the pre-trained reused layers
for layer in reused_layers:
    x = layer(x)

# Add your new, untrained layers
x = tf.keras.layers.Dense(64, activation='relu')(x)
x = tf.keras.layers.Dropout(0.2)(x)
outputs = tf.keras.layers.Dense(your_num_classes, activation='softmax')(x)

new_model = tf.keras.Model(inputs=inputs, outputs=outputs)

Option B: Copy weights to new layers

If you want to tweak reused layers (e.g., rename them, adjust minor settings without changing weight shape), create matching new layers and copy weights over.

# Create new layers matching the structure of the old layers you want to reuse
reused_conv = tf.keras.layers.Conv2D(32, (3,3), activation='relu', name='reused_conv')
reused_pool = tf.keras.layers.MaxPooling2D((2,2), name='reused_pool')

# Copy weights from the old model's corresponding layers
reused_conv.set_weights(old_model.layers[0].get_weights())
reused_pool.set_weights(old_model.layers[1].get_weights())

# Freeze the new reused layers
reused_conv.trainable = False
reused_pool.trainable = False

# Build the new model
inputs = tf.keras.Input(shape=(your_new_input_shape,))
x = reused_conv(inputs)
x = reused_pool(x)

# Add new, trainable layers
x = tf.keras.layers.Flatten()(x)
x = tf.keras.layers.Dense(128, activation='relu')(x)
outputs = tf.keras.layers.Dense(your_num_classes, activation='softmax')(x)

new_model = tf.keras.Model(inputs=inputs, outputs=outputs)

3. Compile and train the new model

Compile the model normally, then start training. Only the non-frozen (new) layers will have their weights updated during training.

new_model.compile(
    optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),
    loss='sparse_categorical_crossentropy',
    metrics=['accuracy']
)

# Train the model — only new layers will be updated
history = new_model.fit(
    train_dataset,
    epochs=15,
    validation_data=val_dataset
)

Key Tips:

  • Shape Alignment: Ensure the input/output shapes of reused layers match the rest of the new model. Mismatched shapes will cause errors.
  • Custom Layers: If your old model uses custom layers, define those same classes in your new code before loading weights.
  • Freezing Timing: Set trainable=False before compiling the model. If you change it post-compile, recompile the model for the setting to take effect.

内容的提问来源于stack exchange,提问作者piratesailor

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.22 08:41:32