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

如何用estimator.export_savedmodel()保存TensorFlow模型及实现serving_input_receiver_fn()?

Alright, let's tackle how to export your custom VGGNet Estimator for TensorFlow Serving using estimator.export_savedmodel()—the tricky part is getting the serving_input_receiver_fn() right, since it needs to align with how your model expects input during inference (matching your training _parse_function but adjusted for serving).

Key Background

The serving_input_receiver_fn defines how your model will accept input from clients when deployed. It needs to:

  • Define placeholder tensors that match the data clients will send (raw image bytes, or preprocessed tensors).
  • Apply the same deterministic preprocessing you used in training (skip random augmentations like random flips/crops—those are only for training!).
  • Map the processed input to the feature names your Estimator's model_fn expects.

Scenario 1: Clients send raw image bytes (JPEG/PNG)

This is the most common case—clients send raw image files, and the serving pipeline handles preprocessing (just like your training _parse_function, minus randomness).

import tensorflow as tf

def serving_input_receiver_fn():
    # Define placeholder to accept raw image bytes (supports batch or single image)
    image_bytes = tf.placeholder(
        dtype=tf.string,
        shape=[None],
        name="client_input_image_bytes"
    )

    # Reuse your training preprocessing logic (remove random augmentations!)
    def preprocess_single_image(byte_str):
        # Match the steps in your _parse_function exactly (adjust for your setup)
        image = tf.image.decode_jpeg(byte_str, channels=3)
        image = tf.image.resize_images(image, [224, 224])  # VGGNet's standard input size
        image = tf.cast(image, tf.float32) / 255.0  # Normalization (match your training step!)
        # Add any other deterministic transforms here (e.g., mean subtraction)
        return image

    # Apply preprocessing to batch of images
    preprocessed_images = tf.map_fn(
        preprocess_single_image,
        image_bytes,
        dtype=tf.float32
    )

    # Return receiver: map processed tensors to your model's input feature name
    return tf.estimator.export.ServingInputReceiver(
        features={"input_layer": preprocessed_images},  # Match your model_fn's input key
        receiver_tensors={"image_bytes": image_bytes}  # Exposed input for clients
    )

Scenario 2: Clients send preprocessed tensors

If clients handle preprocessing themselves (e.g., resized and normalized arrays), the receiver function is simpler:

def serving_input_receiver_fn():
    # Define placeholder matching your model's input shape (VGGNet example)
    input_tensor = tf.placeholder(
        dtype=tf.float32,
        shape=[None, 224, 224, 3],
        name="client_preprocessed_image"
    )

    # Directly map to your model's input feature
    return tf.estimator.export.ServingInputReceiver(
        features={"input_layer": input_tensor},  # Match your model_fn's input key
        receiver_tensors={"preprocessed_image": input_tensor}
    )

Export the Model

Once your serving_input_receiver_fn is ready, export the saved model with:

# Assume your trained estimator is already defined
estimator.export_savedmodel(
    export_dir_base="./vgg_saved_model",  # Directory to save the model
    serving_input_receiver_fn=serving_input_receiver_fn
)

This will create a numbered subdirectory under ./vgg_saved_model (e.g., 1699999999) containing the saved model ready for TensorFlow Serving.

Critical Checks

  1. Preprocessing Consistency: Ensure the deterministic steps in serving_input_receiver_fn match exactly what you did in _parse_function (e.g., same normalization, input size). Mismatches will break inference accuracy.
  2. Feature Name Alignment: The key in features (e.g., "input_layer") must match the feature name your model_fn uses to retrieve input tensors.
  3. Validate the Export: Use TensorFlow's saved_model_cli to verify the input/output nodes:
    saved_model_cli show --dir ./vgg_saved_model/[your_version_number] --all
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:34:32