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

如何在TensorFlow中保存和加载DNN分类器?针对鸢尾花分类示例

Saving and Loading the Iris DNN Classifier in TensorFlow

Got it, let's walk through how to save and load the DNN classifier from the official TensorFlow Iris Estimator example. I'll break this down with code snippets that fit directly into the default program you're working with.

Saving the DNN Classifier

By default, the Estimator API can automatically save your model during training—you just need to tell it where to store the files. Here's how to modify the classifier initialization step to enable saving:

import tensorflow as tf

# Load Iris data (same as the official example)
iris = tf.contrib.learn.datasets.load_dataset('iris')
train_x = iris.data
train_y = iris.target

# Define feature columns (matches the official example)
feature_columns = [tf.feature_column.numeric_column("x", shape=[4])]

# Initialize classifier WITH a model directory specified
classifier = tf.estimator.DNNClassifier(
    feature_columns=feature_columns,
    hidden_units=[10, 10],  # Keep this consistent for later loading
    n_classes=3,            # Matches Iris's 3 flower classes
    model_dir="./iris_trained_model"  # This folder stores all model files
)

# Train the model (same as the official training code)
train_input_fn = tf.estimator.inputs.numpy_input_fn(
    x={"x": train_x},
    y=train_y,
    num_epochs=None,
    shuffle=True
)
classifier.train(input_fn=train_input_fn, steps=2000)

When you run this, TensorFlow will automatically save checkpoints, model graphs, and metadata to the ./iris_trained_model folder during training. It updates these files periodically, so you don’t need to call any extra save functions manually.

Loading the Saved Classifier

To reuse your trained model, you just need to initialize a new DNNClassifier with the exact same configuration as the one you trained, plus the same model_dir. Here's the code:

# Reuse the same feature column definition as training
feature_columns = [tf.feature_column.numeric_column("x", shape=[4])]

# Load the pre-trained model
loaded_classifier = tf.estimator.DNNClassifier(
    feature_columns=feature_columns,
    hidden_units=[10, 10],  # Must match the training setup exactly
    n_classes=3,            # Must match the training class count
    model_dir="./iris_trained_model"  # Point to the saved model folder
)

# Example: Use the loaded model to make predictions
sample_iris_data = [[5.1, 3.5, 1.4, 0.2], [6.2, 3.4, 5.4, 2.3]]
predict_input_fn = tf.estimator.inputs.numpy_input_fn(
    x={"x": sample_iris_data},
    num_epochs=1,
    shuffle=False
)

predictions = list(loaded_classifier.predict(input_fn=predict_input_fn))
for idx, pred in enumerate(predictions):
    print(f"Sample {idx+1}: Predicted class {pred['class_ids'][0]}, Probabilities: {pred['probabilities']}")

Critical Tips

  • Configuration Consistency: The hidden_units, n_classes, and feature_columns must be identical between training and loading. Mismatched parameters will cause TensorFlow to throw errors.
  • Model Folder Integrity: Don’t manually delete or edit files in the model_dir—these include checkpoints, graph definitions, and training metadata that the Estimator needs to reload the model.
  • Adjust Checkpoint Frequency: If you want to save checkpoints less often (default is every 100 steps), add save_checkpoints_steps=500 to the DNNClassifier initialization during training.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 07:04:21