ML-For-Beginners:用 Pickle 序列化 Scikit-learn 模型并构建 Flask 预测 Web 应用
【免费下载链接】ML-For-Beginners12 weeks, 26 lessons, 52 quizzes, classic Machine Learning for all项目地址: https://gitcode.com/GitHub_Trending/ml/ML-For-Beginners
本节是 ML-For-Beginners 课程中的第 3 个应用模块(3-Web-App),聚焦一个典型的“模型落地”场景:将训练好的 Scikit-learn 模型保存为可移植文件,再在一个基于 Flask 的 Web 应用中加载该模型并提供在线预测。你将使用约 8 万条 NUFORC(National UFO Reporting Center)UFO 目击数据训练一个“根据观测秒数、经纬度判断哪个国家报告了 UFO”的模型,掌握pickle序列化、requirements.txt依赖管理,以及 Flask 路由、模板渲染与模型调用的完整链路。学完本篇,你将能独立完成“数据清洗 → 模型训练 → 模型导出 → Web 服务部署”的完整流程。
一、本节目标与整体流程
根据 3-Web-App/README.md 的定义,本节将介绍一个应用型 ML 主题:如何把 Scikit-learn 模型保存为文件,使其能在 Web 应用中用于预测。整体流程包含四个阶段:
- 用 UFO 目击数据清洗、训练一个多分类模型(LogisticRegression);
- 用
pickle把训练好的模型序列化为ufo-model.pkl; - 用 Flask 搭建一个表单驱动的 Web 应用,在
/predict路由中调用模型; - 输入秒数、纬度、经度,返回“最可能报告 UFO 的国家”。
唯一的课程单元是 Build a Web App 教程,配套有空白练习 Notebook(notebook.ipynb)和完整参考答案(solution/notebook.ipynb 与 solution/web-app)。
二、部署考量:模型运行在哪里?
教程指出,构建“消费 ML 模型”的 Web 应用有多种方式,你的 Web 架构甚至可能反过来影响模型的训练方式。假设数据科学团队训练好了模型,你要把它用进应用里,需要先回答这些问题:
- 是 Web 应用还是移动应用?如果是移动应用或 IoT 场景,可以用 TensorFlow Lite 将模型带入 Android/iOS 应用;
- 模型驻留在哪里?云端还是本地?
- 是否需要离线支持?应用必须能离线工作吗?
- 模型用什么技术训练的?训练技术栈会影响部署工具链的选择:
- TensorFlow训练的模型可以用 TensorFlow.js 转换为 Web 应用可用的格式;
- PyTorch训练的模型可以导出为 ONNX(Open Neural Network Exchange)格式,配合 Onnx Runtime 在 JavaScript Web 应用中使用(课程后续会针对 Scikit-learn 模型探索这一选项,可在 4-Classification/4-Applied 看到 ONNX 相关实践);
- Lobe.ai / Azure Custom Vision 等 ML SaaS平台则提供导出模型到多个平台的能力,包括构建一个可被云端查询的定制 API。
此外,还可以构建一个完整的 Flask Web 应用,让模型直接在 Web 浏览器中训练(例如借助 TensorFlow.js 的 JavaScript 环境)。
本节的选择:由于课程一直使用 Python Notebook,因此聚焦“从 Python Notebook 导出训练好的模型,转成 Python Web 应用可读的格式”这一条路径。
三、两个核心工具:Flask 与 Pickle
本节任务只需要两个 Python 生态工具:
Flask
Flask 被其作者定义为"micro-framework"(微型框架),它用 Python 提供 Web 框架的基础功能,并内置模板引擎用于构建网页。它是本应用的服务端框架。
Pickle
Pickle 是 Python 标准库中用于序列化(serialize)与反序列化(de-serialize)Python 对象结构的模块。当你"pickle"一个模型时,实际上是把它拍平(flatten)为字节流,以便在 Web 环境中加载使用。教程特别提醒:pickle 本身不具备安全性,如果别人让你"un-pickle"一个来路不明的文件,务必小心——反序列化任意 pickle 文件可能执行恶意代码。序列化后的模型文件使用.pkl后缀。
四、数据准备:从 8 万条 UFO 记录到干净数据集
数据集来自 NUFORC,共约 8 万条 UFO 目击记录(当前仓库中 data/ufos.csv 实际包含 80,332 行数据)。数据中有意思的描述包括长描述"A man emerges from a beam of light that shines on a grassy field at night and he runs towards the Texas Instruments parking lot"和短描述"the lights chased us"。
CSV 的完整表头为:
datetime,city,state,country,shape,duration (seconds),duration (hours/min),comments,date posted,latitude,longitude即包含目击发生的城市(city)、州(state)、国家(country)、物体形状(shape)以及纬度(latitude)、经度(longitude)等列。在空白 Notebook 中按以下步骤操作:
1. 导入依赖与数据
import pandas as pd import numpy as np ufos = pd.read_csv('./data/ufos.csv') ufos.head()2. 构建精简 DataFrame 并检查标签
ufos = pd.DataFrame({'Seconds': ufos['duration (seconds)'], 'Country': ufos['country'],'Latitude': ufos['latitude'],'Longitude': ufos['longitude']}) ufos.Country.unique()从参考答案 solution/notebook.ipynb 中可以看到,Country字段的唯一值为:# 0 au, 1 ca, 2 de, 3 gb, 4 us,即国家码为澳大利亚(au)、加拿大(ca)、德国(de)、英国(gb)、美国(us)五个取值。
3. 去除空值并限定观测时长
ufos.dropna(inplace=True) ufos = ufos[(ufos['Seconds'] >= 1) & (ufos['Seconds'] <= 60)] ufos.info()只保留 1–60 秒的目击记录,大幅降低需要处理的数据量。
4. 用 LabelEncoder 把国家编码为数字
✅ 提示:LabelEncoder 按字母顺序编码(alphabetically)。
from sklearn.preprocessing import LabelEncoder ufos['Country'] = LabelEncoder().fit_transform(ufos['Country']) ufos.head()清洗后的数据形如(国家码 3 对应 UK,4 对应 US,3 对应 53.2°N/-2.9°E 的英国坐标):
Seconds Country Latitude Longitude 2 20.0 3 53.200000 -2.916667 3 20.0 4 28.978333 -96.645833 14 30.0 4 35.823889 -80.253611 23 60.0 4 45.582778 -122.352222 24 3.0 3 51.783333 -0.783333五、训练模型:LogisticRegression 与约 95% 的准确率
准备好训练/测试划分。选择三个特征作为 X 向量,目标 y 向量为Country——我们的目标是输入Seconds、Latitude、Longitude,返回一个国家 id:
from sklearn.model_selection import train_test_split Selected_features = ['Seconds','Latitude','Longitude'] X = ufos[Selected_features] y = ufos['Country'] X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=0)用逻辑回归训练并评估:
from sklearn.metrics import accuracy_score, classification_report from sklearn.linear_model import LogisticRegression model = LogisticRegression() model.fit(X_train, y_train) predictions = model.predict(X_test) print(classification_report(y_test, predictions)) print('Predicted labels: ', predictions) print('Accuracy: ', accuracy_score(y_test, predictions))准确率约95%,并不令人意外——因为Country与Latitude/Longitude天然强相关。教程坦言:这个模型本身"没什么革命性"(毕竟可以从经纬度推断出国家),但它是一个很好的练习:从原始数据出发,完成清洗、导出,然后把这个模型放进 Web 应用里使用。
六、Pickle 你的模型并本地验证
几行代码即可完成序列化。序列化后,重新加载该模型并用一个包含秒数、纬度、经度的样本数组测试:
import pickle model_filename = 'ufo-model.pkl' pickle.dump(model, open(model_filename,'wb')) model = pickle.load(open('ufo-model.pkl','rb')) print(model.predict([[50,44,-12]]))模型返回'3'——即 UK(英国)的国家编码(gb按字母序排在第 3 位)。
仓库中可直接查看训练好的成品模型文件:solution/ufo-model.pkl。
七、构建 Flask Web 应用
现在构建一个 Flask 应用,以"视觉上更讨喜"的方式调用模型并返回结果。
1. 目录结构
在notebook.ipynb(也就是ufo-model.pkl所在处)旁边创建web-app文件夹,并在其中创建static(内嵌css子文件夹)和templates两个目录:
web-app/ static/ css/ templates/ notebook.ipynb ufo-model.pkl✅ 参考 solution 文件夹查看完成后的应用:solution/web-app。
2. requirements.txt:Python 应用的依赖清单
在 web-app 文件夹中创建requirements.txt。它类似于 JavaScript 应用中的package.json,列出应用所需的全部依赖:
scikit-learn pandas numpy flask然后安装依赖:
cd web-app pip install -r requirements.txt仓库中的实际文件 solution/web-app/requirements.txt 正是这四行内容。
3. styles.css:静态样式
创建static/css/styles.css,包含以下样式(黑底白字、居中网格布局,符合"UFO 夜观"主题):
body { width: 100%; height: 100%; font-family: 'Helvetica'; background: black; color: #fff; text-align: center; letter-spacing: 1.4px; font-size: 30px; } input { min-width: 150px; } .grid { width: 300px; border: 1px solid #2d2d2d; display: grid; justify-content: center; margin: 20px auto; } .box { color: #fff; background: #2d2d2d; padding: 12px; display: inline-block; }对应仓库文件为 solution/web-app/static/css/styles.css。
4. index.html:Jinja2 模板与表单
创建templates/index.html:
<!DOCTYPE html> <html> <head> <meta charset="UTF-8"> <title>🛸 UFO Appearance Prediction! 👽</title> <link rel="stylesheet" href="{{ url_for('static', filename='css/styles.css') }}"> </head> <body> <div class="grid"> <div class="box"> <p>According to the number of seconds, latitude and longitude, which country is likely to have reported seeing a UFO?</p> <form action="{{ url_for('predict')}}" method="post"> <input type="number" name="seconds" placeholder="Seconds" required="required" min="0" max="60" /> <input type="text" name="latitude" placeholder="Latitude" required="required" /> <input type="text" name="longitude" placeholder="Longitude" required="required" /> <button type="submit" class="btn">Predict country where the UFO is seen</button> </form> <p>{{ prediction_text }}</p> </div> </div> </body> </html>对应仓库文件 solution/web-app/templates/index.html。注意模板中的"mustache"语法{{ }}:这些变量(如prediction_text)将由后端应用注入;表单会以 POST 方式提交到/predict路由;静态样式通过url_for('static', filename='css/styles.css')解析。
5. app.py:驱动预测的核心文件
在 web-app 根目录创建app.py:
import numpy as np from flask import Flask, request, render_template import pickle app = Flask(__name__) model = pickle.load(open("./ufo-model.pkl", "rb")) @app.route("/") def home(): return render_template("index.html") @app.route("/predict", methods=["POST"]) def predict(): int_features = [int(x) for x in request.form.values()] final_features = [np.array(int_features)] prediction = model.predict(final_features) output = prediction[0] countries = ["Australia", "Canada", "Germany", "UK", "US"] return render_template( "index.html", prediction_text="Likely country: {}".format(countries[output]) ) if __name__ == "__main__": app.run(debug=True)💡 提示:运行时加
debug=True后,对应用的任何修改都会即时生效,无需重启服务器。注意:生产环境不要开启该模式。
从源码结构看,/predict路由的调用链非常清晰:
request.form.values()收集表单的seconds、latitude、longitude三个值并转为整数列表;np.array(...)包一层构成二维特征数组(model.predict需要 shape 为(n_samples, n_features)的输入,这与训练时X的二维结构一致——模型"记得"自己训练时输入的形状);model.predict(final_features)返回国家编码;- 用
countries列表把编码 0–4 反解回可读国家名(注意该列表顺序必须与 LabelEncoder 的字母序编码一一对应:au→0 Australia, ca→1 Canada, de→2 Germany, gb→3 UK, us→4 US); - 通过
render_template把prediction_text注入index.html重新渲染返回。
仓库中的成品 solution/web-app/app.py 与上述代码一致,唯一区别是模型加载路径写为"../ufo-model.pkl"——因为 solution 目录里web-app/与ufo-model.pkl平级,模型文件在上一级目录。这提示一个实操要点:pickle.load的相对路径是相对于启动app.py时的工作目录/脚本位置解析的,调整目录结构时务必同步修改该路径。
6. 启动应用
运行python app.py或python3 app.py后,本地 Web 服务器启动,打开浏览器填写简短表单,即可回答"UFO 出现在哪个国家"的问题。
八、关键衔接点:输入数据必须与训练时形状一致
教程总结道:用 Flask + 序列化模型消费模型是相当直白的事情,最难的部分是理解"必须向模型发送什么形状的数据才能得到预测"——这完全取决于模型是如何训练的。本模型需要按顺序输入 3 个数据点(Seconds、Latitude、Longitude)。
因此index.html中表单字段的name(seconds/latitude/longitude)及其提交顺序、app.py中request.form.values()的取值顺序、以及训练 Notebook 中Selected_features = ['Seconds','Latitude','Longitude']的顺序,三者必须严格对齐,预测结果才有意义。教程由此引申出一个职业建议:在生产环境中,训练模型的人与在 Web/移动应用消费模型的人之间需要良好沟通——当然,在本课程里这两个人都是你自己。
九、挑战与课后练习
🚀 挑战:教程鼓励你换一种思路——不在 Notebook 里训练再把模型导入 Flask,而是直接在 Flask 应用内部训练模型。尝试把 Notebook 中(数据清洗后)的 Python 代码转换到一个名为train的路由中,让模型在应用内训练,并思考这种做法的利弊(如:模型更新方便,但每次请求/重启的开销、状态管理复杂度等)。
课后作业(assignment.md):用此前回归课程中用过的模型重做这个 Web 应用(例如南瓜数据集),可以保留现有样式,也可以重新设计以呼应南瓜主题;注意把输入字段改成与你所选模型的训练方式相匹配。评分标准:优秀 = 应用如预期运行并部署到云端;合格 = 应用存在缺陷或结果意外;待改进 = 应用无法正常工作。
复习与自研:列出你用 JavaScript 或 Python 构建 Web 应用消费 ML 模型的多种方式,考虑架构问题——模型应留在应用内还是驻留云端?若驻留云端,如何访问?为"应用型 ML Web 方案"画出一张架构模型图。
十、本节文件索引
| 文件 | 说明 |
|---|---|
| 3-Web-App/README.md | 本节总览(本文主体文档) |
| 3-Web-App/1-Web-App/README.md | "Build a Web App" 完整教程 |
| 3-Web-App/1-Web-App/data/ufos.csv | NUFORC UFO 目击数据集(约 8 万行) |
| 3-Web-App/1-Web-App/notebook.ipynb | 空白练习 Notebook |
| 3-Web-App/1-Web-App/solution/notebook.ipynb | 参考答案 Notebook |
| 3-Web-App/1-Web-App/solution/ufo-model.pkl | 序列化的训练模型 |
| 3-Web-App/1-Web-App/solution/web-app/app.py | Flask 应用入口 |
| 3-Web-App/1-Web-App/solution/web-app/requirements.txt | 应用依赖清单 |
| 3-Web-App/1-Web-App/solution/web-app/templates/index.html | 前端表单模板 |
| 3-Web-App/1-Web-App/solution/web-app/static/css/styles.css | 页面样式 |
| 3-Web-App/1-Web-App/assignment.md | 课后作业与评分标准 |
适用前提说明:本模块的 Flask 应用运行在本地(app.run(debug=True)为开发服务器),适用于学习与原型验证;教程本身已提示不要在生产环境开启 debug 模式,若要将该应用部署到云端(作业评分的"优秀"标准之一),需另行考虑生产级 WSGI 服务器与debug=False的配置。
【免费下载链接】ML-For-Beginners12 weeks, 26 lessons, 52 quizzes, classic Machine Learning for all项目地址: https://gitcode.com/GitHub_Trending/ml/ML-For-Beginners
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考