无法保存自定义TensorFlow模型,加载时触发调用函数序列化错误
解决方案
问题出在你定义的Encoder和Decoder是子类化的Keras模型,且使用了save_traces=False保存。这种情况下,SavedModel不会序列化模型的call逻辑,只能通过模型的配置信息重建,而你原代码的模型依赖全局变量、未正确实现序列化所需的get_config方法,导致加载失败。
核心修复步骤
- 把模型初始化依赖的外部参数(
h_size、emb_size、vocab_len)作为__init__的参数传入,避免使用全局变量 - 实现
get_config方法,保存模型初始化所需的所有参数 - 加载模型时指定
custom_objects,告知Keras如何重建自定义模型类
修改后的模型代码
class Encoder(tf.keras.Model): def __init__(self, h_size, emb_size): super().__init__() self.h_size = h_size self.emb_size = emb_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, 'emb_size': self.emb_size }) return config @classmethod def from_config(cls, config): # 从配置重建模型 return cls(**config) class Decoder(tf.keras.Model): def __init__(self, h_size, emb_size, vocab_len): super().__init__() self.h_size = h_size self.emb_size = emb_size self.vocab_len = vocab_len 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({ 'h_size': self.h_size, 'emb_size': self.emb_size, 'vocab_len': self.vocab_len }) return config @classmethod def from_config(cls, config): return cls(**config)
模型初始化与保存
# 定义参数(替换为你的实际值) h_size = 512 emb_size = 64 vocab_len = 10000 random_vector_len = 128 # 初始化模型 encoder_model = Encoder(h_size, emb_size) decoder_model = Decoder(h_size, emb_size, vocab_len) # 必须先构建模型(让Keras确定各层形状) encoder_dummy_input = tf.random.normal((1, 10, random_vector_len)) encoder_model(encoder_dummy_input) decoder_dummy_input = tf.random.uniform((1, 5), maxval=vocab_len, dtype=tf.int32) init_state = encoder_model(encoder_dummy_input) decoder_model(decoder_dummy_input, init_state) # 保存模型 tf.keras.models.save_model(encoder_model, 'encoder', save_traces=False, overwrite=True) tf.keras.models.save_model(decoder_model, 'decoder', save_traces=False, overwrite=True)
模型加载与使用
# 加载时指定自定义模型类 encoder_model = tf.keras.models.load_model('encoder', custom_objects={'Encoder': Encoder}) decoder_model = tf.keras.models.load_model('decoder', custom_objects={'Decoder': Decoder}) # 后续代码正常运行 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) seq2seq.summary() encoder_model.summary() decoder_model.summary()
替代方案(无需修改模型代码)
如果不想调整模型结构,可以将save_traces设为True(默认值),这样SavedModel会保存call函数的计算图,加载时无需指定custom_objects,但模型文件体积会增大:
# 保存时 tf.keras.models.save_model(encoder_model, 'encoder', save_traces=True, overwrite=True) tf.keras.models.save_model(decoder_model, 'decoder', save_traces=True, overwrite=True) # 加载时直接调用 encoder_model = tf.keras.models.load_model('encoder') decoder_model = tf.keras.models.load_model('decoder')
内容的提问来源于stack exchange,提问作者crn4t
相关产品推荐
相关产品推荐

