首页
/ ML-For-Beginners 分类课程实战:用食材数据集完成多分类建模、算法对比与 ONNX 推理 Web 应用

ML-For-Beginners 分类课程实战:用食材数据集完成多分类建模、算法对比与 ONNX 推理 Web 应用

2026-09-06 16:40:20作者:韦蓉瑛

本篇基于 ML-For-Beginners 课程第 4 章《Getting started with classification》(分类入门)展开。该课程以一份覆盖亚洲与印度五种国家料理的食材数据集为载体,带领读者完整走完经典机器学习中"分类(Classification)"任务的全流程:从数据清洗与 SMOTE 类别平衡,到多分类算法选型、Logistic Regression 与多种分类器的精度对比,最后把训练好的模型转换为 ONNX 格式并嵌入一个纯前端的推理 Web 应用。读完本篇,你可以独立复现这套"数据准备 → 模型训练 → 算法对比 → 模型导出与 Web 端推理"的完整实践路径。

泰式街头小吃摊,第 4 章分类课程的主题配图:亚洲与印度料理数据集

课程定位与整体结构

本章的导读文档见 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_classovr(one-vs-rest,为每个类别训练一个二分类器)或 multinomial(交叉熵损失,仅 lbfgssagsaganewton-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,页面由三部分构成:

  1. 复选框组:每个食材复选框的 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>
  1. 引入 Onnx Runtime Web(课程指定版本 1.9.0):
<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web@1.9.0/dist/ort.min.js"></script>
  1. 推理逻辑:构造 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 直接推理,形成无后端依赖的推荐应用原型。

这套流程的价值在于:它把"分类"从抽象概念变成了一条可复制、可验证、可部署的具体链路——同样的方法可以迁移到任何"标签 + 高维稀疏特征"的多分类场景。

登录后查看全文
热门项目推荐
相关项目推荐