如何使用Pickle序列化模型并将其从Flask后端分离?
如何用Pickle分离Flask后端与机器学习训练代码
前辈的建议非常合理——把训练代码和后端服务彻底分离,不仅能大幅提升后端的响应速度,还能避免每次请求都重复执行训练逻辑造成的资源浪费。我来一步步帮你实现这个需求,同时解释Pickle序列化的具体用法。
第一步:单独编写训练脚本,用Pickle保存模型与向量器
首先,我们把所有训练相关的逻辑抽出来,写成一个独立的Python脚本(比如叫train_model.py)。这个脚本只需要运行一次,完成训练后会把训练好的模型和TF-IDF向量器保存成Pickle文件,供后端调用。
import pandas as pd import joblib from sklearn.naive_bayes import MultinomialNB from sklearn.feature_extraction.text import TfidfVectorizer from sklearn.model_selection import train_test_split # 1. 加载并预处理数据集 df = pd.read_csv("YoutubeSpamMergedData.csv") df_data = df[["CONTENT", "CLASS"]] df_x = df_data['CONTENT'] df_y = df_data.CLASS # 2. 初始化并拟合TF-IDF向量器(必须保存,预测时要用到同一个向量器) cv = TfidfVectorizer(ngram_range=[1,2]) X = cv.fit_transform(df_x) # 3. 拆分数据集并训练朴素贝叶斯模型 X_train, X_test, y_train, y_test = train_test_split(X, df_y, test_size=0.33, random_state=42) clf = MultinomialNB() clf.fit(X_train, y_train) # 4. 验证模型效果(可选,打印准确率) acc = clf.score(X_test, y_test) print(f"训练完成!模型准确率: {acc:.2f}") # 5. 用joblib(Pickle的优化版,适合scikit-learn对象)保存向量器和模型 joblib.dump(cv, 'tfidf_vectorizer.pkl') joblib.dump(clf, 'naivebayes_spam_model.pkl') # 可选:把准确率保存到文件,方便后端读取 with open('model_accuracy.txt', 'w') as f: f.write(str(acc))
运行这个脚本后,你会得到三个文件:
tfidf_vectorizer.pkl:保存了训练时用的TF-IDF向量器,预测时必须用同一个向量器转换输入,否则特征不匹配naivebayes_spam_model.pkl:保存了训练好的朴素贝叶斯分类器model_accuracy.txt:保存了模型的准确率(可选)
第二步:修改Flask后端代码,加载Pickle文件做预测
现在把原Flask代码里的训练逻辑全部删掉,只保留Web服务和预测的核心逻辑,启动时直接加载保存好的Pickle文件:
from flask import Flask, render_template, request import joblib app = Flask(__name__) # 启动Flask时一次性加载向量器和模型(只加载一次,提升性能) try: cv = joblib.load('tfidf_vectorizer.pkl') clf = joblib.load('naivebayes_spam_model.pkl') # 读取保存的准确率(如果之前保存了的话) with open('model_accuracy.txt', 'r') as f: ACCURACY = float(f.read()) except FileNotFoundError: raise Exception("请先运行train_model.py生成Pickle文件!") @app.route('/') def home(): return render_template('home.html') @app.route('/predict', methods=['POST']) def predict(): if request.method == 'POST': comment = request.form['comment'] # 用加载好的向量器转换输入评论 data = [comment] vect = cv.transform(data).toarray() # 用加载好的模型做预测 my_prediction = clf.predict(vect) return render_template('result.html', prediction=my_prediction, accuracy=ACCURACY) if __name__ == '__main__': app.run(debug=True)
关键修改说明:
- 移除训练代码:所有数据加载、模型训练的逻辑都移到了独立脚本中,后端只负责加载和使用模型
- 提前加载模型:在Flask启动时就加载向量器和模型,避免每次请求都重复加载,大幅提升响应速度
- 修复原代码的小问题:
- 原代码中缺少
TfidfVectorizer的导入,现在已补充 - 弃用了过时的
from sklearn.externals import joblib,改用直接import joblib - 修正了拼写错误
_name_为__name__
- 原代码中缺少
关于Pickle序列化的补充说明
Pickle是Python内置的序列化工具,它能把Python对象(比如scikit-learn的模型、自定义类实例)转换成字节流保存到文件,之后可以通过pickle.load()或joblib.load()把对象还原回来。对于scikit-learn的模型,推荐用joblib而不是原生pickle,因为joblib对NumPy数组和scikit-learn的对象做了优化,保存的文件更小,加载速度更快。
注意:Pickle文件只能在Python环境中使用,而且不要加载来自不可信来源的Pickle文件,因为它可能包含恶意代码。
内容的提问来源于stack exchange,提问作者Akhil Alexander
相关产品推荐
相关产品推荐

