机器学习部署入门: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$ 的稳定服务要求。

更多推荐