机器学习部署入门:Sklearn 模型序列化与 Flask 接口封装
·
机器学习部署入门:Sklearn 模型序列化与 Flask 接口封装
1. 模型序列化原理
机器学习模型本质是参数化函数,其数学表示为: $$ f(\mathbf{X}) = \mathbf{W}^T \phi(\mathbf{X}) + b $$ 其中 $\mathbf{W}$ 是权重矩阵,$\phi$ 是特征变换函数。序列化通过二进制编码保存模型参数状态。
2. Sklearn 模型序列化
使用 joblib 高效保存/加载模型:
from sklearn.ensemble import RandomForestClassifier
from joblib import dump, load
# 训练模型
model = RandomForestClassifier(n_estimators=100)
model.fit(X_train, y_train)
# 序列化保存 (生成model.joblib文件)
dump(model, 'model.joblib')
# 反序列化加载
loaded_model = load('model.joblib')
3. Flask 接口封装
创建预测 API 服务:
from flask import Flask, request, jsonify
import numpy as np
app = Flask(__name__)
model = load('model.joblib') # 加载序列化模型
@app.route('/predict', methods=['POST'])
def predict():
# 解析JSON输入
data = request.json['features']
# 转换格式并预测
features = np.array(data).reshape(1, -1)
prediction = model.predict(features)
# 返回JSON格式结果
return jsonify({'prediction': int(prediction[0])})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=5000)
4. 接口测试
使用 curl 测试 API:
curl -X POST http://localhost:5000/predict \
-H "Content-Type: application/json" \
-d '{"features": [5.1, 3.5, 1.4, 0.2]}'
返回结果示例:
{"prediction": 0}
5. 部署优化建议
- 输入验证:添加特征维度检查
- 错误处理:捕获预测异常
- 性能优化:
- 使用
gunicorn部署 - 添加缓存机制
- 使用
- 安全措施:
- 请求速率限制
- HTTPS 加密传输
关键点:序列化保证模型参数完整性,RESTful API 实现跨平台调用,满足 $\frac{\partial \text{Deployment}}{\partial t} \to 0$ 的稳定服务要求。
更多推荐
所有评论(0)