TensorFlow:恢复同架构模型并以权重平均值初始化新模型
Great question! Averaging weights across multiple trained models (with identical architectures) is a fantastic way to boost generalization and reduce overfitting—this is often called "weight averaging" and is a common trick in ensemble learning. Your approach is on the right track, but let's refine it to fix small issues and make it more robust.
First, Fix the Typos in Your Example
In your averaged weight/bias dictionaries (w and b), you reference w3 and b3, but these keys don't exist in your original model weight/bias sets. That would throw an error, so we'll remove those in the corrected code.
A Clean, Workable Implementation
Here's a revised version of your code that properly loads the two models, computes weight averages, and initializes a new model with those averages:
import tensorflow as tf import numpy as np # Define a shared function to create model weights/biases # Ensures all models have identical architecture and variable names def build_model_weights(): weights = { 'w1': tf.Variable(tf.random_normal([2, 3]), name='w1'), 'w2': tf.Variable(tf.random_normal([3, 1]), name='w2') } biases = { 'b1': tf.Variable(tf.random_normal([3]), name='b1'), 'b2': tf.Variable(tf.random_normal([1]), name='b2') } return weights, biases # Create weights for model 1, model 2, and our target averaged model weights1, biases1 = build_model_weights() weights2, biases2 = build_model_weights() avg_weights, avg_biases = build_model_weights() # Savers to load your pre-trained models (matches checkpoint variable names) saver1 = tf.train.Saver({'w1': weights1['w1'], 'w2': weights1['w2'], 'b1': biases1['b1'], 'b2': biases1['b2']}) saver2 = tf.train.Saver({'w1': weights2['w1'], 'w2': weights2['w2'], 'b1': biases2['b1'], 'b2': biases2['b2']}) # Define operations to compute and assign averaged weights assign_ops = [] # Average weights for key in weights1.keys(): avg_val = (weights1[key] + weights2[key]) / 2.0 assign_ops.append(tf.assign(avg_weights[key], avg_val)) # Average biases for key in biases1.keys(): avg_val = (biases1[key] + biases2[key]) / 2.0 assign_ops.append(tf.assign(avg_biases[key], avg_val)) # Saver to save our new averaged model avg_saver = tf.train.Saver({'w1': avg_weights['w1'], 'w2': avg_weights['w2'], 'b1': avg_biases['b1'], 'b2': avg_biases['b2']}) # Run the session to execute the process with tf.Session() as sess: # Initialize all variables first sess.run(tf.global_variables_initializer()) # Restore your two pre-trained models saver1.restore(sess, 'my_model1/model_weights.ckpt') saver2.restore(sess, 'my_model2/model_weights.ckpt') # Apply the average to our new model sess.run(assign_ops) # Optional: Verify the average is correct w1_avg = sess.run(avg_weights['w1']) w1_model1 = sess.run(weights1['w1']) w1_model2 = sess.run(weights2['w1']) print("Average computed correctly?", np.allclose(w1_avg, (w1_model1 + w1_model2)/2)) # Save the averaged model for future use save_path = avg_saver.save(sess, 'averaged_model/avg_weights.ckpt') print(f"Averaged model saved to {save_path}")
Key Improvements:
- Shared Architecture Function: Using a function to build weights ensures all models have identical shapes and variable names—critical for compatibility.
- Explicit Assignment: Instead of creating variables with averages upfront, we compute the average of restored weights and assign them to the new model. This avoids graph node mismatches.
- Verification Step: A quick check to confirm the average is calculated properly.
- Saver for Averaged Model: Saves the new model so you can load it later without re-computing averages.
For More Than Two Models
If you have more than two models, extend this by loading all weights into a list and computing the mean:
# Example for 3 models all_weights = [weights1, weights2, weights3] for key in avg_weights.keys(): weight_stack = tf.stack([w[key] for w in all_weights]) avg_val = tf.reduce_mean(weight_stack, axis=0) assign_ops.append(tf.assign(avg_weights[key], avg_val))
This works perfectly because your models share the same architecture—their weights are directly comparable, so averaging them makes mathematical sense. It's a simple way to get better performance without extra training!
内容的提问来源于stack exchange,提问作者dvarelas

