加载ResNet50预训练权重时出现形状不匹配错误,求排查方案
问题:ResNet50权重加载时出现形状不匹配错误
代码片段
第一部分:权重下载代码
import requests url = 'https://github.com/fchollet/deep-learning-models/releases/download/v0.2/resnet50_weights_tf_dim_ordering_tf_kernels_notop.h5' response = requests.get(url) with open('resnet50_weights_tf_dim_ordering_tf_kernels_notop.h5', 'wb') as f: f.write(response.content)
第二部分:模型定义代码
def create_model(input_shape, n_out): input_tensor = Input(shape=input_shape) base_model = applications.ResNet50(weights='imagenet', include_top=False, input_tensor=input_tensor) base_model.load_weights('resnet50_weights_tf_dim_ordering_tf_kernels_notop.h5') x = GlobalAveragePooling2D()(base_model.output) x = Dropout(0.5)(x) x = Dense(2048, activation='relu')(x) x = Dropout(0.5)(x) final_output = Dense(n_out, activation='softmax', name='final_output')(x) model = Model(input_tensor, final_output) return model
第三部分:模型编译代码
model = create_model(input_shape=(HEIGHT, WIDTH, CANAL), n_out=N_CLASSES) for layer in model.layers: layer.trainable = False for i in range(-5, 0): model.layers[i].trainable = True metric_list = ["accuracy"] optimizer = optimizers.Adam(lr=WARMUP_LEARNING_RATE) model.compile(optimizer=optimizer, loss="categorical_crossentropy", metrics=metric_list) model.summary()
报错信息
ValueError: Cannot assign value to variable ' conv3_block1_0_conv/kernel:0': Shape mismatch.The variable shape (1, 1, 256, 512), and the assigned value shape (512, 128, 1, 1) are incompatible.
尝试将weights='None'改为weights='imagenet'后,错误依然存在,请问原因是什么?
解答
问题出在两个核心点:
- 重复加载权重引发冲突:初始化
ResNet50时设了weights='imagenet',这会自动下载并加载适配当前环境的ImageNet权重;之后你又手动调用load_weights()加载本地文件,两次加载的权重维度顺序不匹配,导致张量形状冲突。 - 权重文件与环境维度配置不兼容:你下载的权重文件名带
tf_dim_ordering,是TensorFlow风格的channels_last维度顺序(即输入形状为(高, 宽, 通道数)),但如果你的Keras/TensorFlow环境默认是channels_first(通道在前,(通道数, 高, 宽)),或者输入张量的维度顺序和权重文件不匹配,就会出现形状不兼容的报错。
具体修复方法:
- 删掉重复的权重加载代码:直接移除
base_model.load_weights('resnet50_weights_tf_dim_ordering_tf_kernels_notop.h5')这一行,weights='imagenet'已经会自动加载正确的预训练权重,不需要手动重复加载。 - 确认输入维度顺序:保证你的
input_shape是(HEIGHT, WIDTH, CANAL)这种channels_last格式,和权重文件的维度要求匹配。 - 强制统一维度配置:如果不确定环境默认配置,可以在代码开头添加以下代码强制使用
channels_last:from keras import backend as K K.set_image_data_format('channels_last')
内容的提问来源于stack exchange,提问作者Nur Laici
相关产品推荐
相关产品推荐

