You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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)

关键修改说明:

  1. 移除训练代码:所有数据加载、模型训练的逻辑都移到了独立脚本中,后端只负责加载和使用模型
  2. 提前加载模型:在Flask启动时就加载向量器和模型,避免每次请求都重复加载,大幅提升响应速度
  3. 修复原代码的小问题:
    • 原代码中缺少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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.28 04:22:29