自定义keras.Model类模型使用save_model无法保存权重的原因
自定义Keras模型保存与加载后权重丢失问题解决方案
问题描述
训练继承自tf.keras.Model的自定义模型时流程正常,但保存为.keras格式后重新加载,模型架构保留但所有层权重为空,抛出断言错误:
AssertionError: dense_2 of loaded model has empty weight list
使用Sequential或函数式API时无此问题,仅自定义类模型实现get_config()后出现该问题。
问题根源
build方法签名不符合规范:Keras要求自定义模型的build方法必须接收input_shape参数,否则模型加载后无法根据输入维度正确初始化层的权重结构。- 模型加载后未完成权重初始化:加载后的模型没有通过输入数据触发
build,导致层实例化后未分配权重空间。
修复代码
修改自定义模型的build方法签名,并确保模型加载后完成初始化:
import tensorflow as tf import numpy as np @tf.keras.utils.register_keras_serializable() class dummyModel(tf.keras.Model): def __init__(self, dense_size, **kwargs): super().__init__(**kwargs) self.dense_size = dense_size # 修正build方法,添加input_shape参数并调用父类build方法 def build(self, input_shape): self.layer1 = tf.keras.layers.Dense(self.dense_size) self.layer2 = tf.keras.layers.Dense(self.dense_size) super().build(input_shape) def call(self, inputs): x = self.layer1(inputs) out = self.layer2(x) return out def get_config(self): config = super().get_config() config.update({'dense_size': self.dense_size}) return config # 生成测试数据 input_data = tf.random.normal((100,100,1)) output_data = tf.random.normal((100,100,1)) # 训练模型 model = dummyModel(10) model.compile(optimizer='adam', loss='mse') model.fit(input_data, output_data, epochs=10, batch_size=10) # 保存模型 tf.keras.models.save_model(model, 'dummy.keras') # 加载模型 modelSaved = tf.keras.models.load_model('dummy.keras') # 关键:通过输入数据触发加载后模型的build流程,完成权重初始化 _ = modelSaved(input_data) # 验证架构与权重 assert model.to_json() == modelSaved.to_json(), "Model architectures are different" for layer1, layer2 in zip(model.layers, modelSaved.layers): weights1 = layer1.get_weights() weights2 = layer2.get_weights() if weights1 != []: assert weights2, f"{layer2.name} of loaded model has empty weight list" for w1, w2 in zip(weights1, weights2): # 改用allclose处理浮点精度误差 assert np.allclose(w1, w2), f"Weights of {layer1.name} are different." print("Models are identical (architecture and weights).")
修复要点说明
- 规范
build方法:添加input_shape参数并调用父类build方法,让Keras能正确追踪层的权重结构。 - 触发加载后初始化:通过传入输入数据调用加载后的模型,触发
build流程,完成权重的重建与加载。 - 浮点精度兼容:将
np.array_equal替换为np.allclose,避免训练过程中浮点精度差异导致断言失败。
内容的提问来源于stack exchange,提问作者Moroder
相关产品推荐
相关产品推荐

