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

TensorFlow 2.0加载Estimator导出的SavedModel后如何预测

适配SavedModel加载与预测到你的DDoS分类场景

Hey there, let's break down how to adapt that basic prediction example to your specific setup with the DDoS classifier you saved. I'll walk through each step clearly.

1. First, Recap Your Model Saving Flow

You already did this part correctly—here's your code formatted cleanly for reference:

import tensorflow as tf

# Define your feature columns
feature_columns = [
    tf.feature_column.numeric_column(key='Fwd_IAT_Total', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Flow_Duration', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_Packet_Length_Std', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Init_Win_bytes_forward', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Destination_Port', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Protocol', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_Packet_Length_Min', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Min_Packet_Length', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_Packets/s', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_IAT_Max', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Average_Packet_Size', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_Header_Length', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_Packet_Length_Max', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_Header_Length.1', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Flow_IAT_Min', dtype=tf.float32),
    tf.feature_column.numeric_column(key='min_seg_size_forward', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Fwd_IAT_Mean', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Max_Packet_Length', dtype=tf.float32),
    tf.feature_column.numeric_column(key='ACK_Flag_Count', dtype=tf.float32),
    tf.feature_column.numeric_column(key='Packet_Length_Std', dtype=tf.float32)
]

# Build the serving input receiver function
serving_input_fn = tf.estimator.export.build_parsing_serving_input_receiver_fn(
    tf.feature_column.make_parse_example_spec(feature_columns))

# Export the trained model
estimator_path = classifier.export_saved_model("/model1", serving_input_fn)

2. What You Already Got Right

You successfully loaded the model with tf.saved_model.load("/model1") and got an AutoTrackable object—this means the model loaded properly. The next step is to use the model's prediction signature and format your input data correctly.

3. Adapt the Prediction Example to Your Multi-Feature Model

The basic example only handles a single feature "x", but your model uses 20 different features. Here's how to modify the code to work with your setup:

Full Adapted Prediction Code

import tensorflow as tf

# Load your saved model
ddos_classifier_1 = tf.saved_model.load("/model1")

# Get the prediction signature (check available signatures if "predict" doesn't work)
# Uncomment the line below to see all available signatures:
# print(list(ddos_classifier_1.signatures.keys()))
predict_signature = ddos_classifier_1.signatures["predict"]

def predict_ddos(input_sample):
    """
    Args:
        input_sample: A dictionary where keys match your feature column names,
                      and values are the numeric values for a single sample.
    
    Returns:
        Model predictions (usually includes 'probabilities' and 'classes' keys)
    """
    # Create a tf.train.Example to hold the input data
    example = tf.train.Example()
    
    # Add each feature's value to the Example
    for feature_name, value in input_sample.items():
        example.features.feature[feature_name].float_list.value.append(value)
    
    # Serialize the Example and pass it to the model
    serialized_example = example.SerializeToString()
    prediction_results = predict_signature(examples=tf.constant([serialized_example]))
    
    return prediction_results

How to Use This Function

Here's an example of calling the function with a test sample:

# Example test sample (fill in with your actual data)
test_data = {
    'Fwd_IAT_Total': 120.7,
    'Flow_Duration': 3500.2,
    'Fwd_Packet_Length_Std': 18.5,
    'Init_Win_bytes_forward': 65535.0,
    'Destination_Port': 443.0,
    'Protocol': 6.0,
    'Fwd_Packet_Length_Min': 40.0,
    'Min_Packet_Length': 40.0,
    'Fwd_Packets/s': 15.3,
    'Fwd_IAT_Max': 60.1,
    'Average_Packet_Size': 75.2,
    'Fwd_Header_Length': 20.0,
    'Fwd_Packet_Length_Max': 150.0,
    'Fwd_Header_Length.1': 20.0,
    'Flow_IAT_Min': 12.3,
    'min_seg_size_forward': 40.0,
    'Fwd_IAT_Mean': 30.5,
    'Max_Packet_Length': 150.0,
    'ACK_Flag_Count': 1.0,
    'Packet_Length_Std': 25.4
}

# Get predictions
results = predict_ddos(test_data)

# Print the results (adjust based on your model's output)
print("Predicted probabilities:", results['probabilities'].numpy())
print("Predicted class:", results['classes'].numpy())

Quick Tips

  • If predict isn't a valid signature, check what's available with print(list(ddos_classifier_1.signatures.keys()))—common alternatives are serving_default.
  • Make sure your input_sample includes all 20 feature columns—missing any will cause an error.
  • To predict on multiple samples at once, serialize multiple tf.train.Example objects and pass them as a list to tf.constant().

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 09:27:44