Keras自定义编解码层显示未构建,导致模型保存失败
问题修复方案
核心问题分析
- 自定义层
SimpleEncoder和SimpleDecoder的get_config中手动序列化子层,导致Keras序列化逻辑冲突,无法正确识别层的参数状态。 - 模型
RetrosynthesisSeq2SeqModel的build方法通过调用self.call触发子层构建,这种方式无法正确同步层的构建状态标记,导致summary显示unbuilt。
具体修复步骤
1. 简化自定义层的序列化逻辑
移除手动序列化子层的代码,Keras会自动处理内部可训练层的序列化与反序列化:
修改SimpleEncoder的get_config和from_config
def get_config(self) -> dict: config = super(SimpleEncoder, self).get_config() config.update({ 'vocab_size': self.vocab_size, 'embedding_dim': self.embedding_dim, 'units': self.units, 'dropout_rate': self.dropout_rate, }) return config @classmethod def from_config(cls, config: dict) -> 'SimpleEncoder': return cls(**config)
修改SimpleDecoder的get_config和from_config
def get_config(self) -> dict: config = super(SimpleDecoder, self).get_config() config.update({ 'vocab_size': self.vocab_size, 'embedding_dim': self.embedding_dim, 'units': self.units, 'dropout_rate': self.dropout_rate, }) return config @classmethod def from_config(cls, config: dict) -> 'SimpleDecoder': return cls(**config)
2. 修正模型的build方法
直接调用子层的build方法,确保层状态正确标记为已构建:
def build(self, input_shape): encoder_input_shape, decoder_input_shape = input_shape # 构建编码器 self.encoder.build(encoder_input_shape) # 构建解码器的完整输入形状(包含初始状态形状) decoder_full_input_shape = ( decoder_input_shape, (encoder_input_shape[0], self.units), (encoder_input_shape[0], self.units) ) self.decoder.build(decoder_full_input_shape) # 构建状态转换Dense层 self.enc_state_h.build((encoder_input_shape[0], self.units)) self.enc_state_c.build((encoder_input_shape[0], self.units)) super(RetrosynthesisSeq2SeqModel, self).build(input_shape)
3. 修正模型get_config的参数获取方式
改用初始化参数而非直接访问子层内部属性:
def get_config(self) -> dict: config = super(RetrosynthesisSeq2SeqModel, self).get_config() config.update({ 'units': self.units, 'input_vocab_size': self.input_vocab_size, 'output_vocab_size': self.output_vocab_size, 'encoder_embedding_dim': self.encoder.embedding_dim, 'decoder_embedding_dim': self.decoder.embedding_dim, 'dropout_rate': self.dropout_rate, }) return config
验证修复效果
修改后重新运行测试脚本:
model.summary()会正确显示SimpleEncoder和SimpleDecoder的参数数量。model.save()可以成功保存模型,无序列化错误。
内容的提问来源于stack exchange,提问作者Biggles-2
相关产品推荐
相关产品推荐

