如何在TensorFlow中保存和加载DNN分类器?针对鸢尾花分类示例
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, andfeature_columnsmust 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=500to theDNNClassifierinitialization during training.
内容的提问来源于stack exchange,提问作者Abhisek Roy

