TensorFlow加载模型报错:无法反序列化Generator类如何解决?
解决自定义子类模型保存后无需导入原类即可加载的问题
方法1:转换为Functional格式保存
子类模型可通过输入示例导出为Functional模型,保存后加载无需依赖原类定义:
- 先创建匹配Generator输入形状的张量示例
- 用
tf.keras.models.Model包装成Functional模型再保存
代码示例:
import tensorflow as tf from networks import Generator # 仅保存时需导入一次 # 初始化原模型 generator = Generator(max_conv_dim=512, sn=False, img_size=256, img_ch=3, style_dim=64) # 定义输入张量(batch_size可设为None或具体值) input_img = tf.keras.Input(shape=(256, 256, 3)) input_style = tf.keras.Input(shape=(64,)) # 构建Functional模型 functional_generator = tf.keras.Model( inputs=[input_img, input_style], outputs=generator([input_img, input_style]) ) # 保存模型 functional_generator.save('generator_functional') # 加载时无需导入Generator类 loaded_model = tf.keras.models.load_model('generator_functional') # 直接推理 output = loaded_model([test_image_tensor, test_style_tensor])
方法2:用SavedModel格式导出推理签名
利用SavedModel的签名机制,加载后可直接通过签名调用模型,无需依赖原类:
import tensorflow as tf from networks import Generator # 仅保存时需导入一次 # 初始化并构建模型(需先调用一次确保权重初始化) generator = Generator(max_conv_dim=512, sn=False, img_size=256, img_ch=3, style_dim=64) dummy_img = tf.random.normal((1, 256, 256, 3)) dummy_style = tf.random.normal((1, 64)) _ = generator([dummy_img, dummy_style]) # 保存为SavedModel格式 generator.save('generator_savedmodel') # 加载时无需导入原类 loaded_generator = tf.saved_model.load('generator_savedmodel') # 获取推理签名(可通过print(loaded_generator.signatures)查看签名名) infer_fn = loaded_generator.signatures['serving_default'] # 执行推理(输入需匹配签名的张量名) output = infer_fn(input_1=dummy_img, input_2=dummy_style)['output_0']
注意事项
- 保存前必须确保模型已被构建(至少传入一次符合形状的输入),否则Keras无法序列化模型结构
- 上述两种方法保存的模型,加载后无需依赖原
Generator类和ops.py中的自定义层,层结构已被完整序列化到保存文件中
内容的提问来源于stack exchange,提问作者let me down slowly
相关产品推荐
相关产品推荐

