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

自定义keras.Model类模型使用save_model无法保存权重的原因

自定义Keras模型保存与加载后权重丢失问题解决方案

问题描述

训练继承自tf.keras.Model的自定义模型时流程正常,但保存为.keras格式后重新加载,模型架构保留但所有层权重为空,抛出断言错误:

AssertionError: dense_2 of loaded model has empty weight list

使用Sequential或函数式API时无此问题,仅自定义类模型实现get_config()后出现该问题。

问题根源

  1. build方法签名不符合规范:Keras要求自定义模型的build方法必须接收input_shape参数,否则模型加载后无法根据输入维度正确初始化层的权重结构。
  2. 模型加载后未完成权重初始化:加载后的模型没有通过输入数据触发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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 15:08:24