自定义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
相关产品推荐
相关产品推荐

