You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将自定义预处理函数嵌入tf.Keras模型图作为SavedModel签名?

解决方法:将预处理函数作为SavedModel签名嵌入tf.Keras模型

你的核心问题在于原代码的两个关键错误:一是在@tf.function中调用model.predict()(这会触发Eager模式推理,无法嵌入静态图);二是签名函数将模型作为参数传入,导致保存后无法与模型实例绑定。以下是两种可行的实现方案:

方案一:用tf.Module封装模型与预处理逻辑

将模型和预处理逻辑封装到tf.Module中,定义绑定到实例的签名方法,确保预处理和模型推理都被序列化到SavedModel的计算图里。

import tensorflow as tf

# 假设已训练好的Keras基础模型
base_model = tf.keras.applications.ResNet50(weights='imagenet')

class InferenceModule(tf.Module):
    def __init__(self, model):
        self.model = model
    
    # 明确指定输入签名,确保SavedModel能识别输入输出格式
    @tf.function(input_signature=[tf.TensorSpec(shape=(), dtype=tf.string)])
    def process(self, img_path):
        # 所有预处理逻辑必须用TensorFlow原生操作实现(不能用非TF库直接处理)
        # 读取图片文件
        img_raw = tf.io.read_file(img_path)
        # 解码为RGB图像
        img = tf.image.decode_jpeg(img_raw, channels=3)
        # 缩放至模型要求的尺寸
        img = tf.image.resize(img, (224, 224))
        # 模型专属预处理
        img = tf.keras.applications.resnet50.preprocess_input(img)
        # 添加batch维度
        img = tf.expand_dims(img, axis=0)
        
        # 直接调用模型(替代model.predict(),确保嵌入静态图)
        preds = self.model(img)
        return {"preds": preds}

# 创建模块实例
inference_module = InferenceModule(base_model)

# 保存模型并指定自定义签名
tf.saved_model.save(inference_module, "./saved_model", signatures={"process": inference_module.process})

# 加载测试示例
loaded_module = tf.saved_model.load("./saved_model")
test_img = tf.constant("test.jpg")
result = loaded_module.process(test_img)
print(result["preds"].shape)

方案二:自定义Keras模型整合预处理

如果你需要保留Keras模型的原生结构,可以自定义模型类,将预处理逻辑整合到call方法,或单独绑定签名函数。

import tensorflow as tf

class CustomInferenceModel(tf.keras.Model):
    def __init__(self, base_model):
        super().__init__()
        self.base_model = base_model
    
    def call(self, img_path):
        # 预处理逻辑与方案一一致,全部用TF操作实现
        img_raw = tf.io.read_file(img_path)
        img = tf.image.decode_jpeg(img_raw, channels=3)
        img = tf.image.resize(img, (224, 224))
        img = tf.keras.applications.resnet50.preprocess_input(img)
        img = tf.expand_dims(img, axis=0)
        return self.base_model(img)

# 初始化自定义模型
base_model = tf.keras.applications.ResNet50(weights='imagenet')
custom_model = CustomInferenceModel(base_model)

# 定义绑定到模型实例的签名函数
@tf.function(input_signature=[tf.TensorSpec(shape=(), dtype=tf.string)])
def process_fn(img_path):
    return {"preds": custom_model(img_path)}

# 保存模型并指定签名
tf.saved_model.save(custom_model, "./saved_keras_model", signatures={"process": process_fn})

# 加载测试示例
loaded_model = tf.saved_model.load("./saved_keras_model")
test_img = tf.constant("test.jpg")
result = loaded_model.process(test_img)
print(result["preds"].shape)

关键注意事项

  • 预处理必须用TF原生操作:如果使用非TensorFlow库(如PIL)处理图像,需要用tf.py_function包装,但会降低性能且可能导致跨环境兼容性问题,优先用TF原生API实现。
  • 避免在@tf.function中调用predict():model.predict()是Eager模式的高层API,无法嵌入静态图,必须直接用model(input_tensor)的方式调用。
  • 明确输入签名:通过input_signature指定输入张量的形状和类型,确保SavedModel能正确生成可复用的签名。

内容的提问来源于stack exchange,提问作者BenyaminGhN

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.31 09:31:01