You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

TensorFlow Keras使用load_weights跳过不兼容层为何仍报错?

问题:tf.keras load_weights设置skip_mismatch=True仍报形状不匹配错误

尝试使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.12 14:06:03