TensorFlow自定义Encoder/Decoder模型无法保存问题求助
解决TensorFlow自定义模型无法保存的问题
问题根源
你的自定义Encoder和Decoder模型无法保存的核心问题在于**get_config方法实现错误**:
Decoder类中super().get_config(self)调用错误,正确写法是不带self参数的super().get_config()get_config直接返回了层对象(比如self.lstm1),而TensorFlow无法序列化层实例,必须返回可序列化的层配置字典- 模型初始化依赖的外部参数(
h_size、vocab_len、emb_size)未纳入配置,导致重建模型时无法还原结构
解决方案1:正确实现get_config和from_config
修正后的模型代码确保配置可序列化且能正确重建模型:
修正后的Encoder类
class Encoder(tf.keras.Model): def __init__(self, h_size): super().__init__() self.h_size = h_size self.lstm1 = tf.keras.layers.LSTM(h_size * 2, return_sequences=False, return_state=True, name='lstm1_enc') self.lstm2 = tf.keras.layers.LSTM(h_size * 2, return_sequences=True, return_state=True, name='lstm2_enc') self.lstm3 = tf.keras.layers.LSTM(h_size * 4, return_sequences=True, return_state=True, name='lstm3_enc') def call(self, x): state = [] out, h, c = self.lstm3(x) state.append((h, c)) out, h, c = self.lstm2(out) state.append((h, c)) out, h, c = self.lstm1(out) state.append((h, c)) return state def get_config(self): config = super().get_config() config.update({ 'h_size': self.h_size, 'lstm1': self.lstm1.get_config(), 'lstm2': self.lstm2.get_config(), 'lstm3': self.lstm3.get_config() }) return config @classmethod def from_config(cls, config): model = cls(config['h_size']) model.lstm1 = tf.keras.layers.LSTM.from_config(config['lstm1']) model.lstm2 = tf.keras.layers.LSTM.from_config(config['lstm2']) model.lstm3 = tf.keras.layers.LSTM.from_config(config['lstm3']) return model
修正后的Decoder类
class Decoder(tf.keras.Model): def __init__(self, vocab_len, emb_size, h_size): super().__init__() self.vocab_len = vocab_len self.emb_size = emb_size self.h_size = h_size self.embed = tf.keras.layers.Embedding(vocab_len, emb_size, name='embed') self.lstm1 = tf.keras.layers.LSTM(h_size * 2, return_sequences=True, return_state=True, name='lstm1_dec') self.lstm2 = tf.keras.layers.LSTM(h_size * 2, return_sequences=True, return_state=True, name='lstm2_dec') self.lstm3 = tf.keras.layers.LSTM(h_size * 4, return_sequences=True, return_state=True, name='lstm3_dec') self.fc = tf.keras.layers.Dense(vocab_len, activation='softmax', name='fc') def call(self, x, init_state): state = [] out = self.embed(x) out, h, c = self.lstm1(out, initial_state=init_state[2]) state.append((h, c)) out, h, c = self.lstm2(out, initial_state=init_state[1]) state.append((h, c)) out, h, c = self.lstm3(out, initial_state=init_state[0]) state.append((h, c)) out = self.fc(out) return out, state def get_config(self): config = super().get_config() config.update({ 'vocab_len': self.vocab_len, 'emb_size': self.emb_size, 'h_size': self.h_size, 'embed': self.embed.get_config(), 'lstm1': self.lstm1.get_config(), 'lstm2': self.lstm2.get_config(), 'lstm3': self.lstm3.get_config(), 'fc': self.fc.get_config() }) return config @classmethod def from_config(cls, config): model = cls(config['vocab_len'], config['emb_size'], config['h_size']) model.embed = tf.keras.layers.Embedding.from_config(config['embed']) model.lstm1 = tf.keras.layers.LSTM.from_config(config['lstm1']) model.lstm2 = tf.keras.layers.LSTM.from_config(config['lstm2']) model.lstm3 = tf.keras.layers.LSTM.from_config(config['lstm3']) model.fc = tf.keras.layers.Dense.from_config(config['fc']) return model
初始化与保存
传入实际参数初始化模型,之后即可正常保存:
# 替换为你实际的参数值 h_size = 64 vocab_len = 1000 emb_size = 128 random_vector_len = 32 encoder_model = Encoder(h_size) decoder_model = Decoder(vocab_len, emb_size, h_size) # 保持原seq2seq构建代码不变 encoder_inputs = tf.keras.layers.Input(shape=(None, random_vector_len)) decoder_inputs = tf.keras.layers.Input(shape=(None, )) enc_state = encoder_model(encoder_inputs) decoder_outputs, _ = decoder_model(decoder_inputs, enc_state) seq2seq = tf.keras.Model([encoder_inputs, decoder_inputs], decoder_outputs) # 执行保存 tf.keras.models.save_model(encoder_model, 'encoder1', overwrite=True) tf.keras.models.save_model(decoder_model, 'decoder', overwrite=True)
解决方案2:仅保存权重(快速替代方案)
若不想修改get_config,可使用save_weights仅保存权重,后续重建模型结构后加载:
保存权重
encoder_model.save_weights('encoder_weights.h5') decoder_model.save_weights('decoder_weights.h5')
加载权重
# 重建Encoder并加载权重 encoder_model = Encoder(h_size) dummy_enc_input = tf.random.normal((1, 10, random_vector_len)) encoder_model(dummy_enc_input) encoder_model.load_weights('encoder_weights.h5') # 重建Decoder并加载权重 decoder_model = Decoder(vocab_len, emb_size, h_size) dummy_dec_input = tf.random.uniform((1, 5), maxval=vocab_len, dtype=tf.int32) dummy_state = [ (tf.random.normal((1, h_size*4)), tf.random.normal((1, h_size*4))), (tf.random.normal((1, h_size*2)), tf.random.normal((1, h_size*2))), (tf.random.normal((1, h_size*2)), tf.random.normal((1, h_size*2))) ] decoder_model(dummy_dec_input, dummy_state) decoder_model.load_weights('decoder_weights.h5')
内容的提问来源于stack exchange,提问作者crn4t
相关产品推荐
相关产品推荐

