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

自定义Keras模型保存后加载失败:形状不匹配问题求助

问题分析与解决方法

核心原因

加载时出现形状不匹配,本质是模型序列化/反序列化过程中,共享层的变量映射出错——HighNet复用的BaseNet编码器/解码器层,在保存时没有被正确记录状态,加载时Keras无法匹配到对应形状的权重,导致错误赋值。

具体解决步骤

1. 固定层的初始化逻辑

所有层必须在模型的__init__方法中定义,禁止在call方法里动态创建层。动态创建层会导致Keras无法捕获层的完整形状信息,保存时丢失关键配置。

示例修正后的BaseNet:

import tensorflow as tf
from tensorflow import keras

class BaseNet(keras.Model):
    def __init__(self, latent_dim=24):
        super().__init__()
        self.latent_dim = latent_dim
        # 所有层在__init__中初始化,固定输入输出形状
        self.encoder = keras.Sequential([
            keras.layers.Flatten(input_shape=(28, 28)),
            keras.layers.Dense(256, activation='relu'),
            keras.layers.Dense(latent_dim * 2)  # 输出均值+方差
        ])
        self.decoder = keras.Sequential([
            keras.layers.Dense(256, activation='relu', input_shape=(latent_dim,)),
            keras.layers.Dense(784, activation='sigmoid'),
            keras.layers.Reshape((28, 28))
        ])
    
    def call(self, x):
        z_mean, z_log_var = tf.split(self.encoder(x), num_or_size_splits=2, axis=1)
        z = self.reparameterize(z_mean, z_log_var)
        return self.decoder(z), z_mean, z_log_var
    
    def reparameterize(self, mean, log_var):
        eps = tf.random.normal(shape=mean.shape)
        return mean + tf.exp(0.5 * log_var) * eps
    
    # 重写get_config,确保自定义参数被保存
    def get_config(self):
        config = super().get_config()
        config.update({'latent_dim': self.latent_dim})
        return config

2. 正确复用已训练层

HighNet必须直接接收BaseNet训练好的encoder和decoder实例,不能重新定义相同结构的层。同时实现自定义模型的序列化方法,确保共享层的状态被正确保存。

示例修正后的HighNet:

class HighNet(keras.Model):
    def __init__(self, encoder, decoder):
        super().__init__()
        self.encoder = encoder
        self.decoder = decoder
        # HighNet自有层同样在__init__中初始化
        self.classifier = keras.Sequential([
            keras.layers.Dense(128, activation='relu', input_shape=(encoder.output_shape[1]//2,)),
            keras.layers.Dense(10, activation='softmax')
        ])
    
    def call(self, x):
        z_mean, z_log_var = tf.split(self.encoder(x), num_or_size_splits=2, axis=1)
        z = self.reparameterize(z_mean, z_log_var)
        recon = self.decoder(z)
        cls_out = self.classifier(z_mean)
        return recon, cls_out, z_mean, z_log_var
    
    def reparameterize(self, mean, log_var):
        eps = tf.random.normal(shape=mean.shape)
        return mean + tf.exp(0.5 * log_var) * eps
    
    # 自定义序列化逻辑,保存编码器和解码器的配置
    def get_config(self):
        config = super().get_config()
        config.update({
            'encoder': keras.layers.serialize(self.encoder),
            'decoder': keras.layers.serialize(self.decoder)
        })
        return config
    
    # 自定义反序列化逻辑,重建编码器和解码器实例
    @classmethod
    def from_config(cls, config):
        encoder = keras.layers.deserialize(config.pop('encoder'))
        decoder = keras.layers.deserialize(config.pop('decoder'))
        return cls(encoder, decoder, **config)

3. 标准保存与加载流程

使用Keras原生的save和load_model方法,加载时必须传入所有自定义模型类到custom_objects参数,避免Keras无法识别模型结构:

# 保存HighNet模型
high_net.save('high_net_model.keras')

# 加载模型
loaded_high_net = keras.saving.load_model(
    'high_net_model.keras',
    custom_objects={'HighNet': HighNet, 'BaseNet': BaseNet}
)

4. 额外检查点

  • 确保BaseNet训练完成后,再将其encoder和decoder传入HighNet,避免层的形状未固定。
  • 不要混用pickle/dill保存Keras模型,这类工具无法正确处理Keras层的内部变量依赖关系。

内容的提问来源于stack exchange,提问作者grimbar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 03:22:56