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

import joblib
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

# Load the saved model
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

import pickle
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

# Use pickle to load the saved model
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.pipeline import Pipeline
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

# Load the pipeline
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

import time

# 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

from flask import Flask, request, jsonify
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)
Other Extensions