Keras如何部分加载模型权重?以VGG19的block1为例
如何部分加载VGG19模型的block1权重
这个报错很典型——你直接调用model.load_weights()时,Keras默认会尝试匹配权重文件里所有层的参数,但你自定义的模型只包含VGG19的block1部分,和权重文件的层数量、结构不匹配,所以就抛出了这个ValueError。
下面给你两种可行的解决方案:
方法一:使用by_name=True参数(最简单)
Keras的load_weights()方法支持by_name=True参数,它会根据层的名字去匹配加载权重,只会加载那些在你的自定义模型中存在对应名字的层的参数,完全忽略其他层。
修改你的代码如下:
import tensorflow as tf def VGG19_part(input_shape=None): img_input = tf.keras.layers.Input(shape=input_shape) # Block 1 x = tf.keras.layers.Conv2D(64, (3, 3), activation='linear', padding='same', name='block1_conv1')(img_input) x = tf.keras.layers.Activation('relu')(x) x = tf.keras.layers.Conv2D(64, (3, 3), activation='linear', padding='same', name='block1_conv2')(x) x = tf.keras.layers.Activation('relu')(x) x = tf.keras.layers.MaxPooling2D((2, 2), strides=(2, 2), name='block1_pool')(x) model = tf.keras.Model(img_input, x, name='vgg19') # 关键修改:添加by_name=True参数 model.load_weights('/Users/myuser/.keras/models/vgg19_weights_tf_dim_ordering_tf_kernels_notop.h5', by_name=True) print(model.summary()) # 可以验证权重是否加载成功,比如打印block1_conv1的权重形状 print("block1_conv1权重形状:", model.get_layer('block1_conv1').get_weights()[0].shape) return model # 测试调用 model = VGG19_part(input_shape=(224,224,3))
这样就能成功加载block1的所有层权重了,因为你的自定义模型里的层名字和权重文件里的block1层名字完全一致(block1_conv1、block1_conv2、block1_pool),Keras会自动匹配这些层并加载对应的权重。
方法二:手动提取权重(更灵活)
如果你需要更精细的控制,可以先加载完整的VGG19模型,然后手动把block1的层权重提取出来,赋值给你的自定义模型:
import tensorflow as tf def VGG19_part(input_shape=None): img_input = tf.keras.layers.Input(shape=input_shape) # Block 1 x = tf.keras.layers.Conv2D(64, (3, 3), activation='linear', padding='same', name='block1_conv1')(img_input) x = tf.keras.layers.Activation('relu')(x) x = tf.keras.layers.Conv2D(64, (3, 3), activation='linear', padding='same', name='block1_conv2')(x) x = tf.keras.layers.Activation('relu')(x) x = tf.keras.layers.MaxPooling2D((2, 2), strides=(2, 2), name='block1_pool')(x) model = tf.keras.Model(img_input, x, name='vgg19') # 加载完整的VGG19模型(不带顶层) full_vgg = tf.keras.applications.VGG19(weights='imagenet', include_top=False, input_shape=input_shape) # 遍历自定义模型的层,从完整模型中提取对应权重 for layer in model.layers: if layer.name in [l.name for l in full_vgg.layers]: full_layer = full_vgg.get_layer(layer.name) layer.set_weights(full_layer.get_weights()) print(f"已加载层 {layer.name} 的权重") print(model.summary()) return model # 测试调用 model = VGG19_part(input_shape=(224,224,3))
这种方法适合你需要对权重做额外处理(比如修改、筛选)的场景,灵活性更高。
注意点
- 确保你自定义模型的层名字和权重文件中的层名字完全一致,否则
by_name=True或者手动提取都会匹配失败。 - 如果你使用的是带顶层的权重文件(比如
vgg19_weights_tf_dim_ordering_tf_kernels.h5),同样可以用这两种方法,因为顶层的层名字不在你的自定义模型中,会被自动忽略。
内容的提问来源于stack exchange,提问作者mrgloom
相关产品推荐
相关产品推荐

