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=Falsebefore compiling the model. If you change it post-compile, recompile the model for the setting to take effect.
内容的提问来源于stack exchange,提问作者piratesailor

