Sklearn Model Saving and Loading
In machine learning, the model training process is usually time-consuming. To avoid retraining the model every time, we can save the trained model for later loading and prediction.
scikit-learnThere are two common ways to save and load models:joblibandpickle。
1. UsingjoblibSave and Load Models
joblibis an efficient Python serialization tool, especially suitable for saving objects containing large amounts of numerical arrays (such as numpy arrays, scikit-learn models, etc.). Compared topickle,joblibit is more efficient when processing large-scale data.
joblib is an external Python library that can be installed with the following command:
pip install joblib
Save Model
joblib provides a simple API to save and load objects.
We can use the joblib.dump() method to save the model to a file.
Example
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.svm import SVC
# Load data
data = load_iris()
X, y = data.data, data.target
# Split training set and test set
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
# Create and train the model
model = SVC(kernel='linear')
model.fit(X_train, y_train)
# Save model to file
joblib.dump(model, 'svm_model.joblib')
Load Model
Use the joblib.load() method to load the saved model object.
Example
loaded_model = joblib.load('svm_model.joblib')
# Use the loaded model for prediction
y_pred = loaded_model.predict(X_test)
# Print the prediction results
print("Predictions:", y_pred)
Through the above steps, we successfully saved the trained model to a file and can load the model and make predictions at any later time.
2. Using pickle to Save and Load Models
pickle is a built-in Python module that allows Python objects to be serialized and deserialized.
Although joblib is more suitable for handling large amounts of data, pickle is also a common tool for saving and loading models, suitable for general cases.
Save Model
Similar to joblib, pickle also has a simple API to save and load objects.
The code to save the model is as follows:
Example
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.svm import SVC
# Load data
data = load_iris()
X, y = data.data, data.target
# Split training set and test set
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
# Create and train the model
model = SVC(kernel='linear')
model.fit(X_train, y_train)
# Save the model using pickle
with open('svm_model.pkl', 'wb') as f:
pickle.dump(model, f)
Load Model
Use pickle.load() to load the model:
Example
with open('svm_model.pkl', 'rb') as f:
loaded_model = pickle.load(f)
# Use the loaded model for prediction
y_pred = loaded_model.predict(X_test)
# Print the prediction results
print("Predictions:", y_pred)
3、joblib vs pickle
joblib and pickle are two common methods for saving and loading models.
joblib is more suitable for saving large data objects, while pickle is Python's standard serialization tool, suitable for general cases.
joblib: Usually suitable for saving objects containing large amounts of numerical data (such as numpy arrays).joblibWhen processing large-scale data, it is more efficient thanpicklemore efficient.pickle: Suitable for saving smaller objects or regular Python objects. It is a built-in Python library and does not require additional installation.
If the model contains a large number of numerical arrays or matrices (such as support vector machines, random forests, etc.), joblib is recommended because it is more efficient than pickle. For smaller models or models that do not contain large amounts of numerical data, pickle is sufficient.
4. Save and Load Pipeline
In practical applications, a model is not just a single model; sometimes it combines multiple processing steps (such as data preprocessing, feature selection, model training, etc.). These processing steps can be accomplished using scikit-learn's Pipeline. The Pipeline can also be saved and loaded using joblib or pickle.
Save Pipeline:
Example
from sklearn.preprocessing import StandardScaler
from sklearn.svm import SVC
import joblib
# Create a pipeline
pipeline = Pipeline([
('scaler', StandardScaler()),
('svc', SVC(kernel='linear'))
])
# Train the pipeline
pipeline.fit(X_train, y_train)
# Save the pipeline to a file
joblib.dump(pipeline, 'pipeline_model.joblib')
Load Pipeline:
Example
loaded_pipeline = joblib.load('pipeline_model.joblib')
# Use the loaded pipeline for prediction
y_pred = loaded_pipeline.predict(X_test)
# Print the prediction results
print("Predictions:", y_pred)
The process of saving and loading a pipeline is the same as that of a single model; you just need to ensure that the entire pipeline object is saved and loaded.
5. Model Version Management
In real-world machine learning applications, model updates and version management are crucial. Each time you train and save a model, it is best to add a timestamp or version number to the model file name to distinguish between different versions of the model. For example:
Example
# Create a timestamp
timestamp = time.strftime("%Y%m%d-%H%M%S")
# Save the model with a timestamp
joblib.dump(model, f'svm_model_{timestamp}.joblib')
In this way, we can manage different versions of models based on timestamps, making it easier to roll back and update models.
6. Using the Model for Persistence
Once the model is trained and saved, we can load it in subsequent practical applications to make predictions without retraining.
For example, we can integrate the saved model with web services, batch jobs, or other applications, so that the model can be reused without retraining.
Using Loaded Models in Web Services
For example, suppose we are using Flask to create a simple web service that provides model prediction services through an API. In this case, we can load the saved model for real-time prediction.
Example
import joblib
import numpy as np
app = Flask(__name__)
# Load the model
model = joblib.load('svm_model.joblib')
@app.route('/predict', methods=['POST'])
def predict():
data = request.get_json() # Get input data
features = np.array(data['features']).reshape(1, -1) # Convert to a format suitable for prediction
prediction = model.predict(features) # Use the loaded model for prediction
return jsonify({'prediction': prediction.tolist()}) # Return the prediction result
if __name__ == '__main__':
app.run(debug=True)