本文目录导读:

- 使用 pickle(最通用)
- 使用 joblib(推荐用于scikit-learn)
- 使用 ONNX(跨平台兼容)
- 使用 TensorFlow/Keras
- 使用 PyTorch
- 保存额外信息
- 完整的保存/加载示例
- 推荐选择:
- 注意事项:
在Python中保存和加载机器学习模型有多种方法,以下是几种常见的方式:
使用 pickle(最通用)
import pickle
from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import make_classification
# 训练模型
X, y = make_classification(n_samples=1000, n_features=4)
model = RandomForestClassifier()
model.fit(X, y)
# 保存模型
with open('model.pkl', 'wb') as f:
pickle.dump(model, f)
# 加载模型
with open('model.pkl', 'rb') as f:
loaded_model = pickle.load(f)
# 使用模型进行预测
predictions = loaded_model.predict(X)
使用 joblib(推荐用于scikit-learn)
import joblib
from sklearn.ensemble import RandomForestClassifier
# 训练模型
model = RandomForestClassifier()
model.fit(X, y)
# 保存模型
joblib.dump(model, 'model.joblib')
# 加载模型
loaded_model = joblib.load('model.joblib')
# 使用模型
predictions = loaded_model.predict(X)
使用 ONNX(跨平台兼容)
import onnx
import skl2onnx
from skl2onnx import convert_sklearn
from skl2onnx.common.data_types import FloatTensorType
# 训练模型
from sklearn.linear_model import LogisticRegression
model = LogisticRegression()
model.fit(X, y)
# 转换为ONNX格式
initial_type = [('float_input', FloatTensorType([None, X.shape[1]]))]
onnx_model = convert_sklearn(model, initial_types=initial_type)
# 保存ONNX模型
with open("model.onnx", "wb") as f:
f.write(onnx_model.SerializeToString())
# 加载ONNX模型
import onnxruntime
sess = onnxruntime.InferenceSession("model.onnx")
input_name = sess.get_inputs()[0].name
output_name = sess.get_outputs()[0].name
# 预测
result = sess.run([output_name], {input_name: X.astype(np.float32)})[0]
使用 TensorFlow/Keras
import tensorflow as tf
from tensorflow import keras
# 创建和训练模型
model = keras.Sequential([
keras.layers.Dense(64, activation='relu'),
keras.layers.Dense(10, activation='softmax')
])
model.compile(optimizer='adam', loss='sparse_categorical_crossentropy')
model.fit(X, y, epochs=10)
# 保存模型(HDF5格式)
model.save('model.h5')
# 加载模型
loaded_model = keras.models.load_model('model.h5')
# 或保存为SavedModel格式(推荐)
model.save('saved_model')
# 加载SavedModel
loaded_model = tf.keras.models.load_model('saved_model')
使用 PyTorch
import torch
import torch.nn as nn
# 定义模型
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(4, 1)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
# 保存整个模型
torch.save(model, 'model.pth')
# 推荐:只保存状态字典
torch.save(model.state_dict(), 'model_params.pth')
# 加载模型
# 方式1:加载整个模型
loaded_model = torch.load('model.pth')
# 方式2:先创建模型实例,再加载参数
model = SimpleModel()
model.load_state_dict(torch.load('model_params.pth'))
model.eval()
保存额外信息
import json
# 保存模型和元数据
model_info = {
'model': joblib.dump(model, 'model.joblib'),
'feature_names': ['feature1', 'feature2', 'feature3', 'feature4'],
'model_type': 'RandomForest',
'accuracy': 0.95,
'training_date': '2024-01-15'
}
with open('model_info.json', 'w') as f:
json.dump(model_info, f, indent=2)
完整的保存/加载示例
import joblib
import pandas as pd
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
from sklearn.ensemble import RandomForestClassifier
# 完整的模型训练和保存流程
def train_and_save_model():
# 准备数据
data = pd.read_csv('data.csv')
X = data.drop('target', axis=1)
y = data['target']
# 分割数据
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)
# 标准化
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)
# 训练模型
model = RandomForestClassifier()
model.fit(X_train_scaled, y_train)
# 保存模型和scaler
joblib.dump(model, 'model.joblib')
joblib.dump(scaler, 'scaler.joblib')
# 保存特征名称
with open('feature_names.txt', 'w') as f:
f.write('\n'.join(X.columns))
print("模型保存完成!")
# 加载和使用模型
def load_and_predict(new_data):
# 加载模型
model = joblib.load('model.joblib')
scaler = joblib.load('scaler.joblib')
# 加载特征名称
with open('feature_names.txt', 'r') as f:
feature_names = f.read().splitlines()
# 预处理新数据
new_data_scaled = scaler.transform(new_data[feature_names])
# 预测
predictions = model.predict(new_data_scaled)
probabilities = model.predict_proba(new_data_scaled)
return predictions, probabilities
推荐选择:
- scikit-learn模型:使用
joblib - 深度学习模型:使用框架自带的保存方法(Keras的
.h5或PyTorch的.pth) - 跨平台兼容:使用ONNX格式
- 简单模型:使用
pickle
注意事项:
- 保存模型时记住保存数据预处理对象(scaler, encoder等)
- 保持Python版本和依赖库版本一致性
- 大型模型建议使用joblib而不是pickle
- 生产环境考虑使用版本控制和模型注册表