ML-For-Beginners 分类课程实战:用食材数据集完成多分类建模、算法对比与 ONNX 推理 Web 应用
本篇基于 ML-For-Beginners 课程第 4 章《Getting started with classification》(分类入门)展开。该课程以一份覆盖亚洲与印度五种国家料理的食材数据集为载体,带领读者完整走完经典机器学习中"分类(Classification)"任务的全流程:从数据清洗与 SMOTE 类别平衡,到多分类算法选型、Logistic Regression 与多种分类器的精度对比,最后把训练好的模型转换为 ONNX 格式并嵌入一个纯前端的推理 Web 应用。读完本篇,你可以独立复现这套"数据准备 → 模型训练 → 算法对比 → 模型导出与 Web 端推理"的完整实践路径。
课程定位与整体结构
本章的导读文档见 4-Classification/README.md,其核心设定是:
- 区域主题:美味的亚洲与印度料理。课程提出一个有趣的问题——观察一组食材,能否判断出这道菜属于哪种菜系?
- 学习路径:在回归(Regression)章节的基础上继续学习,掌握可用于"更好地理解数据"的其他分类器;
- 配套数据:食材数据集来自 Kaggle 的 Asian and Indian Cuisines 数据集,存放于 4-Classification/data/cuisines.csv。
章节包含四课,每课均提供 README 文档、可运行的 notebook.ipynb、练习与课后作业(assignment.md):
| 课 | 标题 | 核心内容 | 文档 |
|---|---|---|---|
| 1 | Introduction to classification | 分类的概念、数据清洗、SMOTE 平衡 | 1-Introduction/README.md |
| 2 | More classifiers | 算法选型逻辑、Logistic Regression 多分类 | 2-Classifiers-1/README.md |
| 3 | Yet other classifiers | Linear SVC、KNN、SVC、Random Forest、AdaBoost 对比 | 3-Classifiers-2/README.md |
| 4 | Applied ML: build a web app | 模型转 ONNX、Netron 查看、Web 端推理 | 4-Applied/README.md |
从课程结构看,这是一条典型的"教学递进"路径:先做数据工程(第 1 课),再建立"如何选算法"的判断框架(第 2、3 课),最后进入工程落地(第 4 课)。
分类任务与数据集
什么是分类
分类(Classification)是一种监督学习形式,与回归高度相似:回归预测连续数值(比如"9 月 vs 12 月的南瓜价格"),逻辑回归预测二值类别("这个南瓜是不是橙色"),而分类则回答"这个数据点属于哪一类"——从简单的二分类(邮件是否垃圾)到复杂的多类问题(图像分割)都适用。其本质是建立一个预测模型,刻画输入变量到输出类别之间的映射。
本课程的食材数据正是一个多分类问题:给定一批食材(特征),判断它属于哪种国家菜系(类别)。原始数据 cuisines.csv 的规模为 2448 行 × 385 列,共 5 个菜系类别,380 个食材特征列均为 0/1 二值(表示某道菜谱是否含该食材)。用命令可核对类别分布:
awk -F',' 'NR>1 {c[$1]++} END {for (k in c) print k, c[k]}' data/cuisines.csv
# korean 799
# indian 598
# chinese 442
# japanese 320
# thai 289
类别之间明显不均衡(korean 是 thai 的 2.8 倍)——这正是第 1 课要解决的数据问题。
第 1 课:数据清洗与 SMOTE 平衡
第 1 课的目标是产出一份干净、平衡、可直接训练的数据集 cleaned_cuisines.csv,供后续三课使用。
环境准备:安装 imblearn
pip install imblearn
imblearn 是 Scikit-learn 的配套扩展包,提供处理类别不均衡的工具(本课使用其中的 SMOTE)。导入所需依赖:
import pandas as pd
import matplotlib.pyplot as plt
import matplotlib as mpl
import numpy as np
from imblearn.over_sampling import SMOTE
读取并检查数据
df = pd.read_csv('../data/cuisines.csv')
df.head() # 前 5 行
df.info() # 2448 行,385 列,int64(384) + object(1)
df.info() 的输出与上文的 awk 统计一致:2448 条菜谱、385 列(384 个数值列 + 1 个 cuisine 字符串标签列)。
观察菜系分布与特征
先看类别分布:
df.cuisine.value_counts().plot.barh()
再按菜系拆分子集并统计:
thai_df = df[(df.cuisine == "thai")] # (289, 385)
japanese_df = df[(df.cuisine == "japanese")] # (320, 385)
chinese_df = df[(df.cuisine == "chinese")] # (442, 385)
indian_df = df[(df.cuisine == "indian")] # (598, 385)
korean_df = df[(df.cuisine == "korean")] # (799, 385)
接下来,用课程提供的 create_ingredient_df() 函数统计每个菜系中出现频率最高的 10 种食材(转置求和 → 过滤零列 → 按频次降序):
def create_ingredient_df(df):
ingredient_df = df.T.drop(['cuisine','Unnamed: 0']).sum(axis=1).to_frame('value')
ingredient_df = ingredient_df[(ingredient_df.T != 0).any()]
ingredient_df = ingredient_df.sort_values(by='value', ascending=False, inplace=False)
return ingredient_df
thai_ingredient_df = create_ingredient_df(thai_df)
thai_ingredient_df.head(10).plot.barh()
对照各菜系的 Top 10 食材图可以直观看到:rice(米)、garlic(蒜)、ginger(姜)几乎在每个菜系的高频榜单里都出现——它们是"跨菜系共有特征",对判别菜系贡献有限,反而制造噪声,因此课程要求将其剔除:
feature_df = df.drop(['cuisine','Unnamed: 0','rice','garlic','ginger'], axis=1)
labels_df = df.cuisine
用 SMOTE 平衡数据集
"Synthetic Minority Over-sampling Technique"(合成少数类过采样)通过插值生成新的少数类样本,把每个类别都拉到最大类的样本数(本例为 799),避免模型因为多数类样本更多而偏向预测多数类:
oversample = SMOTE()
transformed_feature_df, transformed_label_df = oversample.fit_resample(feature_df, labels_df)
print(f'new label count: {transformed_label_df.value_counts()}')
# korean 799 / chinese 799 / indian 799 / japanese 799 / thai 799
最后合并标签与特征并导出(注意仓库中已包含产物文件,行数 3995 行 × 382 列):
transformed_df = pd.concat([transformed_label_df, transformed_feature_df], axis=1, join='outer')
transformed_df.head()
transformed_df.info()
transformed_df.to_csv("../data/cleaned_cuisines.csv")
仓库中 4-Classification/data/cleaned_cuisines.csv 即为该课的输出产物,可直接用于后续课程,无需重跑。
第 2 课:算法选型与 Logistic Regression 多分类
选择分类器的判断框架
Scikit-learn 将分类归入 Supervised Learning,可选方法包括线性模型、支持向量机、随机梯度下降、K 近邻、高斯过程、决策树、集成方法(Voting Classifier)与多类/多输出算法等。课程给出的选型推理(结合 Microsoft Algorithm Cheat Sheet 的多分类部分)是:
- 神经网络太重:数据量小且本地 Notebook 训练,排除;
- 不用 one-vs-all 两分类器方案:排除;
- 决策树或逻辑回归可行;
- 多分类 Boosted Decision Trees 面向排序类非参数任务,不适用。
第 3 课还会给出 Scikit-learn 官方 ML Map 上更细化的路径:样本数 >50、预测类别、有标签、样本数 <100K → 首选 Linear SVC;不行再试 KNN、SVC 与集成分类器。
关键参数:multi_class 与 solver
Logistic Regression 原生设计用于二分类,多分类时由两个参数决定行为:
multi_class:ovr(one-vs-rest,为每个类别训练一个二分类器)或multinomial(交叉熵损失,仅lbfgs、sag、saga、newton-cg等 solver 支持);solver:优化问题的求解算法,不同 solver 与multi_class的可用组合不同(Scikit-learn 文档提供了对照表)。
实验:训练与评估
import pandas as pd
cuisines_df = pd.read_csv("../data/cleaned_cuisines.csv")
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split, cross_val_score
from sklearn.metrics import accuracy_score, precision_score, confusion_matrix, classification_report, precision_recall_curve
from sklearn.svm import SVC
import numpy as np
cuisines_label_df = cuisines_df['cuisine']
cuisines_feature_df = cuisines_df.drop(['Unnamed: 0', 'cuisine'], axis=1)
X_train, X_test, y_train, y_test = train_test_split(cuisines_feature_df, cuisines_label_df, test_size=0.3)
使用 ovr 方案 + liblinear solver 训练(注意用 np.ravel() 压平标签):
lr = LogisticRegression(multi_class='ovr', solver='liblinear')
model = lr.fit(X_train, np.ravel(y_train))
accuracy = model.score(X_test, y_test)
print("Accuracy is {}".format(accuracy)) # 约 80%
对测试集中第 50 行(X_test.iloc[50])做单条预测可看到模型给出的概率分解——indian 0.716、chinese 0.229、japanese 0.030、korean 0.017、thai 0.008;用 predict_proba() 排序即可确认"印度菜"是最高概率的猜测。完整的分类报告:
y_pred = model.predict(X_test)
print(classification_report(y_test, y_pred))
| 类别 | precision | recall | f1-score | support |
|---|---|---|---|---|
| chinese | 0.73 | 0.71 | 0.72 | 229 |
| indian | 0.91 | 0.93 | 0.92 | 254 |
| japanese | 0.70 | 0.75 | 0.72 | 220 |
| korean | 0.86 | 0.76 | 0.81 | 242 |
| thai | 0.79 | 0.85 | 0.82 | 254 |
| accuracy | 0.80 | 1199 |
第 3 课:五种分类器横向对比
第 3 课用一个逐步扩展的 classifiers 字典统一训练、预测并打印报告,便于同条件对比:
from sklearn.neighbors import KNeighborsClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.svm import SVC
from sklearn.ensemble import RandomForestClassifier, AdaBoostClassifier
C = 10
classifiers = {
'Linear SVC': SVC(kernel='linear', C=C, probability=True, random_state=0),
'KNN classifier': KNeighborsClassifier(C),
'SVC': SVC(),
'RFST': RandomForestClassifier(n_estimators=100),
'ADA': AdaBoostClassifier(n_estimators=100),
}
n_classifiers = len(classifiers)
for index, (name, classifier) in enumerate(classifiers.items()):
classifier.fit(X_train, np.ravel(y_train))
y_pred = classifier.predict(X_test)
accuracy = accuracy_score(y_test, y_pred)
print("Accuracy (train) for %s: %0.1f%% " % (name, accuracy * 100))
print(classification_report(y_test, y_pred))
各分类器的实验结果(测试集 1199 条):
| 分类器 | 说明 | 测试准确率 |
|---|---|---|
| Linear SVC | SVC(kernel='linear', C=10, probability=True, random_state=0),probability=True 才可获得概率估计 |
78.6% |
| KNN | KNeighborsClassifier(10),按文档输出 |
73.8% |
| SVC(默认 RBF 核) | SVC() 默认参数 |
83.2% |
| Random Forest | RandomForestClassifier(n_estimators=100),"随机化的决策树森林"平均法 |
84.5% |
| AdaBoost | AdaBoostClassifier(n_estimators=100),逐轮聚焦错分样本的加权修正 |
72.4% |
几个值得注意的现象:
- 课程文档中 Linear SVC 的结果输出(78.6%)中,indian 类的 F1 最高(0.87),thai 最低(0.78);
- 默认 RBF 核的 SVC(83.2%)明显优于线性 SVC(78.6%),说明这个 0/1 特征空间中类别边界并非线性可分;
- Random Forest 以 84.5% 成为本数据集上的最佳模型,而 AdaBoost(72.4%)显著偏弱——集成方法的表现高度依赖基学习器与数据特性;
- 所有模型的薄弱点相似:chinese 类普遍 recall 偏低,说明中式菜谱与日式/韩式菜谱在 380 维食材空间中重叠较多。
关于参数的含义,课程原文强调:SVC 中 kernel 决定如何"聚类"标签空间,C 是正则化强度,probability 默认为 False、需显式置 True 才能输出概率;random_state=0 用于在启用概率估计时固定数据洗牌顺序,保证结果可复现。
第 4 课:模型转 ONNX 与 Web 端推理应用
第 4 课把前面学到的分类技术转化为可落地的"推荐系统"雏形:训练一个模型 → 转为 ONNX → 在纯 JavaScript 页面中做推理。仓库内 4-Classification/4-Applied/solution/ 已包含成品三件套:训练 Notebook(notebook.ipynb)、模型文件(model.onnx)与前端页面(index.html),可对照源码逐步复现。
训练分类模型
沿用第 3 课中表现较好的 SVC 参数:
!pip install skl2onnx
import pandas as pd
data = pd.read_csv('../data/cleaned_cuisines.csv')
X = data.iloc[:, 2:] # 去掉前两列(索引与 cuisine 标签),保留 380 个食材特征
y = data[['cuisine']]
from sklearn.model_selection import train_test_split
from sklearn.svm import SVC
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3)
model = SVC(kernel='linear', C=10, probability=True, random_state=0)
model.fit(X_train, y_train.values.ravel())
y_pred = model.predict(X_test)
print(classification_report(y_test, y_pred)) # 整体 accuracy 0.79
转换为 ONNX
转换时必须声明输入张量的形状——该数据集有 380 种食材,所以是 FloatTensorType([None, 380]):
from skl2onnx import convert_sklearn
from skl2onnx.common.data_types import FloatTensorType
initial_type = [('float_input', FloatTensorType([None, 380]))]
options = {id(model): {'nocl': True, 'zipmap': False}}
onx = convert_sklearn(model, initial_types=initial_type, options=options)
with open("./model.onnx", "wb") as f:
f.write(onx.SerializeToString())
两个转换选项值得理解:
zipmap=False:分类模型默认输出"类别 → 概率"的字典列表,这里去掉后直接输出标签数组,前端读取更简单;nocl=True:不把类别字符串信息写进模型,减小模型体积(前端页面代码中results.label.data[0]读取到的就是类别标签)。
输入名 float_input 是后面 Web 代码 feeds 字典的键名,二者必须严格一致;可用 Netron 打开 model.onnx 可视化核对输入、输出与 380 个输入节点。
Web 端推理
在同一目录新建 index.html,页面由三部分构成:
- 复选框组:每个食材复选框的
value是该食材在数据集特征列中的下标(从 0 计)。例如apple按字母序排在第 5 列,所以value="4";pear为 247、sake为 302、soy sauce为 327、cumin为 112。这些下标可以直接对照仓库中的 4-Classification/data/ingredient_indexes.csv(第 1 行食材名、第 2 行对应下标)查询核对:
<div class="boxCont">
<input type="checkbox" value="4" class="checkbox">
<label>apple</label>
</div>
...
<button onClick="startInference()">What kind of cuisine can you make?</button>
- 引入 Onnx Runtime Web(课程指定版本 1.9.0):
<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web@1.9.0/dist/ort.min.js"></script>
- 推理逻辑:构造 380 维全 0 数组,勾选复选框时把对应下标置 1;点击按钮后异步加载
./model.onnx,封装成ort.Tensor送入名为float_input的输入端执行推理:
const ingredients = Array(380).fill(0);
const checks = [...document.querySelectorAll('.checkbox')];
checks.forEach(check => {
check.addEventListener('change', function() {
ingredients[check.value] = check.checked ? 1 : 0;
});
});
async function startInference() {
if (!checks.some(check => check.checked)) {
alert('Please select at least one ingredient.');
return;
}
try {
const session = await ort.InferenceSession.create('./model.onnx');
const input = new ort.Tensor(new Float32Array(ingredients), [1, 380]);
const feeds = { float_input: input };
const results = await session.run(feeds);
alert('You can enjoy ' + results.label.data[0] + ' cuisine today!');
} catch (e) {
console.error(e);
}
}
本地运行
安装 http-server 后,在 index.html 所在目录执行 http-server,打开 localhost 即可交互使用。仓库成品页面 4-Classification/4-Applied/solution/index.html 与文档代码一致,model.onnx 输入端名为 float_input、形状 [1, 380] 的设置可以直接验证。
这条"Python 训练 → ONNX 导出 → 前端推理"的管线,让模型摆脱了 Python 运行时依赖,可以在浏览器(甚至离线环境)中直接使用;而此前第 3 章的 UFO 回归模型采用的是 Flask + pickle 的全栈 Python 方案,两种部署架构在本仓库中都有实例可供对照。
课后练习与延伸
四课各自配有作业文档:1-Introduction/assignment.md 要求在 Scikit-learn 文档中找出分类算法并与本课程的某个数据集、可提出的问题做匹配表;2-Classifiers-1/assignment.md 要求深入研究不同 solver 的底层行为;3-Classifiers-2/assignment.md 要求研究各分类器的默认参数并思考调参对模型质量的影响;4-Applied/assignment.md 则鼓励把最小推荐应用继续扩展,用 ingredient_indexes.csv 中的全部食材探索不同菜系的味道组合。
此外,每课均提供 R 语言的并行实现(solution/ 目录下的 .Rmd 与渲染后的 .html)与英文版 Notebook 解答(solution/*.ipynb),可按语言偏好对照学习。
小结
本章以 2448 条亚洲与印度菜谱、380 维二值食材特征为实验场,演示了经典分类任务的完整工程闭环:
- 数据工程:识别并剔除跨类别共有特征(rice/garlic/ginger),用 SMOTE 把 289~799 条的失衡类别统一到 799 条,产出 cleaned_cuisines.csv;
- 算法框架:从 ML Map 与 Cheat Sheet 出发建立"数据特征 → 候选算法"的推理习惯,而不是盲目试错;
- 多分类细节:
multi_class(ovr/multinomial)与solver的组合约束是 Logistic Regression 多分类落地的关键; - 实测结论:在本数据集上 Random Forest(84.5%)> SVC(83.2%)> Logistic Regression(80%)> Linear SVC(78.6%)> KNN(73.8%)> AdaBoost(72.4%);
- 部署落地:
skl2onnx转换时以[None, 380]声明输入、float_input命名输入端,前端用 Onnx Runtime Web 直接推理,形成无后端依赖的推荐应用原型。
这套流程的价值在于:它把"分类"从抽象概念变成了一条可复制、可验证、可部署的具体链路——同样的方法可以迁移到任何"标签 + 高维稀疏特征"的多分类场景。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0624
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
