TensorFlow Keras使用load_weights跳过不兼容层为何仍报错?
尝试使用tf.keras.Model的load_weights函数,调用model.load_weights(weights_path, by_name=True, skip_mismatch=True)时仍出现形状不匹配错误——这正是期望skip_mismatch参数能处理的场景。以下是基于MNIST数据集的复现代码(运行于Google Colab):
import tensorflow as tf import numpy as np import pandas as pd import matplotlib.pyplot as plt import os # 导入MNIST数据集 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.mnist.load_data() # 构建简单的全连接序列模型 model = tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape=(28,28)), tf.keras.layers.Rescaling(1./255, input_shape=(28,28)), tf.keras.layers.Dense(128, activation='relu'), tf.keras.layers.Dropout(0.2), tf.keras.layers.Dense(10) ]) # 编译模型 model.compile(optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=['accuracy']) # 创建检查点目录 os.makedirs("checkpoints", exist_ok=True) cp_callback = tf.keras.callbacks.ModelCheckpoint("checkpoints/cp-{epoch:04d}.ckpt", save_weights_only=True, ) # 训练模型 history = model.fit(x_train, y_train, epochs=10, validation_data=(x_test, y_test), callbacks=[cp_callback]) # 构建输出层为5个神经元的模型 model2 = tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape=(28,28)), tf.keras.layers.Rescaling(1./255, input_shape=(28,28)), tf.keras.layers.Dense(128, activation='relu'), tf.keras.layers.Dropout(0.2), tf.keras.layers.Dense(5) ]) # 编译model2 model2.compile(optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=['accuracy']) # 加载model的权重到model2 model2.load_weights("checkpoints/cp-0010.ckpt",by_name=True, skip_mismatch=True)
报错信息:
"Received incompatible tensor with shape (10,) when attempting to restore variable with shape (5,) and name dense_5/bias:0."
无论设置skip_mismatch=True还是False都会触发该错误,请问是否用法有误?正确的使用方式是什么?
原因与解决方案
问题根源
skip_mismatch=True的作用是跳过权重文件与目标模型中存在性不匹配的变量(即权重文件有但模型没有,或模型有但权重文件没有的变量),但它不处理名称存在但形状不匹配的情况——这是当前TensorFlow的设计逻辑。
你的两个模型中,最后一层的自动生成名称是相同的(比如dense_1),by_name=True会匹配到该层,此时形状不匹配(一个是10维,一个是5维),即使开启skip_mismatch=True也会报错。
正确用法
有两种方式可以解决这个问题:
1. 给层手动指定唯一名称
修改模型层的名称,让需要跳过的层与权重文件中的层名称不匹配,这样by_name=True会自动跳过加载:
# 第一个模型(训练用)指定明确名称 model = tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape=(28,28)), tf.keras.layers.Rescaling(1./255, input_shape=(28,28)), tf.keras.layers.Dense(128, activation='relu', name='hidden_dense'), tf.keras.layers.Dropout(0.2), tf.keras.layers.Dense(10, name='output_dense_10') ]) # 第二个模型(加载权重用)给输出层指定不同名称 model2 = tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape=(28,28)), tf.keras.layers.Rescaling(1./255, input_shape=(28,28)), tf.keras.layers.Dense(128, activation='relu', name='hidden_dense'), tf.keras.layers.Dropout(0.2), tf.keras.layers.Dense(5, name='output_dense_5') ])
此时加载权重时,hidden_dense的权重会被正确匹配加载,output_dense_5因名称不匹配会被skip_mismatch=True跳过。
2. 手动过滤并加载权重
如果不想修改层名称,可以手动加载权重文件,过滤掉形状不匹配的变量后再赋值给模型:
# 加载权重文件到字典 weight_dict = tf.train.load_checkpoint("checkpoints/cp-0010.ckpt") # 获取model2的可训练变量,以变量名为键存储 model2_vars = {var.name.split(':')[0]: var for var in model2.trainable_variables} # 遍历权重字典,只加载名称和形状都匹配的权重 for weight_name, weight_value in weight_dict.items(): if weight_name in model2_vars and weight_value.shape == model2_vars[weight_name].shape: model2_vars[weight_name].assign(weight_value.numpy())
这种方式更灵活,可完全控制哪些权重被加载。
补充说明
by_name=True是按层的名称匹配权重,而非层在模型中的位置;skip_mismatch=True仅处理变量存在性不匹配的场景,不解决名称匹配但形状冲突的问题。
内容的提问来源于stack exchange,提问作者Crucio

