ML-For-Beginners 分类入门:基于亚洲美食数据集的数据探索、特征筛选与 SMOTE 类别均衡实战
2026/9/6 17:48:10 网站建设 项目流程

ML-For-Beginners 分类入门:基于亚洲美食数据集的数据探索、特征筛选与 SMOTE 类别均衡实战

【免费下载链接】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 课程"Getting started with classification"单元的第 1 课(分类入门)为主体,完整还原了从一个 2448 行、385 列的亚洲美食数据集出发,经历数据加载、类别分布诊断、典型食材挖掘、易混淆特征剔除,到用 SMOTE 合成少数类过采样完成类别均衡的完整数据准备流程。读完后你不仅掌握多分类问题的定义与判断方法,还能获得一份可直接供后续分类器(Logistic 回归、SVM 等)使用的清洗后数据集cleaned_cuisines.csv

一、分类是什么:与回归的关系及两类基本形态

分类(Classification)是经典机器学习中的核心监督学习任务。用更科学的表述说:你的分类方法要构建一个预测模型,建立输入变量到输出变量之间的映射关系,从而把数据点划分到不同的类别中。分类问题一般分为两大形态:

  • 二分类(binary classification):输出只有两个类别,例如"这封邮件是不是垃圾邮件""这个南瓜是不是橙色"。
  • 多分类(multiclass classification):输出有若干个互斥类别,例如根据一组食材判断它属于哪种菜系。

原文档通过与前面课程的概念对照帮助定位分类在整个课程体系中的位置:

技术解决的问题例子
线性回归预测变量间关系,估计新数据点在关系线上的取值预测 9 月和 12 月南瓜的价格
逻辑回归发现"二分类"边界在这个价格点上,这个南瓜是不是橙色
分类算法族用多种算法判定数据点所属的标签或类别根据一组食材推断菜系来源

分类同样带有监督学习的标签机制:数据是带标签的,算法利用标签学习"特征 → 类别"的映射,与统计学习中的统计分类一脉相承。典型的应用如利用smokerweightage等特征判断"患 X 疾病的可能性"。

本单元四节课统一使用同一个数据集来贯穿整个流程:一份覆盖亚洲与印度五大菜系(thai、japanese、chinese、indian、korean)的食材数据。核心问题是一个多分类问题:给定一批食材,它最可能属于哪一种菜系?原始数据位于 cuisines.csv,本课时完成的清洗、均衡后的产物 cleaned_cuisines.csv 将直接供 第 2 课"使用更多分类器" 使用。

二、实战步骤 1:安装依赖并加载数据

开始动手前,先完成数据清理与**均衡(balance)**这两件准备工作。课程提供了一份空白的 notebook.ipynb 作为练习载体,完整可运行的参考实现在 solution/notebook.ipynb。

2.1 安装 imblearn

第一步安装 imbalanced-learn 包(文中写作 imblearn),这是一个 Scikit-learn 风格的扩展库,提供了后续要用的 SMOTE 过采样等类别均衡工具:

pip install imblearn

2.2 导入依赖

导入数据读取与可视化的库,并从imblearn中导入SMOTE

import pandas as pd import matplotlib.pyplot as plt import matplotlib as mpl import numpy as np from imblearn.over_sampling import SMOTE

2.3 读取数据

使用read_csv()读取原始数据集(相对路径以练习 notebook 所在目录为起点,即本单元目录的上一级data文件夹):

df = pd.read_csv('../data/cuisines.csv')

2.4 检查数据形状

df.head()

前五行数据的结构如下(385 列中大部分为 0/1 的食材特征列):

| | Unnamed: 0 | cuisine | almond | angelica | anise | anise_seed | apple | apple_brandy | apricot | armagnac | ... | whiskey | white_bread | white_wine | whole_grain_wheat_flour | wine | wood | yam | yeast | yogurt | zucchini | | --- | ---------- | ------- | ------ | -------- | ----- | ---------- | ----- | ------------ | ------- | -------- | --- | ------- | ----------- | ---------- | ----------------------- | ---- | ---- | --- | ----- | ------ | -------- | | 0 | 65 | indian | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | ... | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | | 1 | 66 | indian | 1 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | ... | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | | 2 | 67 | indian | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | ... | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | | 3 | 68 | indian | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | ... | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | | 4 | 69 | indian | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | ... | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 0 | 1 | 0 |

2.5 查看数据基本信息

df.info()

输出如下,说明数据共 2448 行、385 列,除cuisine列是object类型外,其余 384 列都是int64的 0/1 食材指示列:

<class 'pandas.core.frame.DataFrame'> RangeIndex: 2448 entries, 0 to 2447 Columns: 385 entries, Unnamed: 0 to zucchini dtypes: int64(384), object(1) memory usage: 7.2+ MB

对原始 CSV 文件实际检查后可以确认:文件共 2449 行(1 行表头 + 2448 行数据),首列Unnamed: 0是无意义的行号(65 起),第二列cuisine是标签列,其后为almondangelicaanise……直至zucchini的 383 个食材特征列。仓库中还附带了一份 ingredient_indexes.csv,记录了 380 个食材名称与其 0 起索引的映射关系,便于理解特征列名与位置索引的对应。

三、实战步骤 2:诊断类别分布

在清理、可视化、准备 ML 任务之前,首先要搞清楚数据在每个菜系上的分布情况。

3.1 横向条形图观察分布

df.cuisine.value_counts().plot.barh()

菜系的种类是有限的,但各菜系的样本量分布明显不均匀——这是后续要修复的核心问题。

3.2 按菜系切分并统计样本数

thai_df = df[(df.cuisine == "thai")] japanese_df = df[(df.cuisine == "japanese")] chinese_df = df[(df.cuisine == "chinese")] indian_df = df[(df.cuisine == "indian")] korean_df = df[(df.cuisine == "korean")] print(f'thai df: {thai_df.shape}') print(f'japanese df: {japanese_df.shape}') print(f'chinese df: {chinese_df.shape}') print(f'indian df: {indian_df.shape}') print(f'korean df: {korean_df.shape}')

输出(与仓库中实际数据一致):

thai df: (289, 385) japanese df: (320, 385) chinese df: (442, 385) indian df: (598, 385) korean df: (799, 385)

最少的 thai(289 条)与最多的 korean(799 条)之间相差近 2.8 倍。这种偏斜会直接影响分类器:如果大多数样本都属于某一个类,模型会倾向于"随大流"地多预测那一类,仅仅因为该类的数据更多。

四、实战步骤 3:挖掘每个菜系的典型食材

接下来深入数据,弄清各菜系的典型食材有哪些。目标之一是剔除在各菜系中反复出现、造成类间混淆的特征——这类"人人爱吃"的食材对区分菜系没有帮助。

4.1 编写 create_ingredient_df 函数

该函数先丢弃无用的列,再按出现次数统计各食材:

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

从源码结构看,其逻辑分三步:先把df转置使食材成为行,drop(['cuisine','Unnamed: 0'])去掉标签列与行号列;sum(axis=1)沿轴 1 求和即得到每个食材在该菜系样本中出现的总次数,to_frame('value')收成单列;第二行(ingredient_df.T != 0).any()过滤掉计数为 0 的食材(即该菜系根本不用到的食材);最后按value降序排列。

4.2 逐个菜系绘制 Top 10 食材

泰国菜:

thai_ingredient_df = create_ingredient_df(thai_df) thai_ingredient_df.head(10).plot.barh()

日本菜:

japanese_ingredient_df = create_ingredient_df(japanese_df) japanese_ingredient_df.head(10).plot.barh()

中国菜:

chinese_ingredient_df = create_ingredient_df(chinese_df) chinese_ingredient_df.head(10).plot.barh()

印度菜:

indian_ingredient_df = create_ingredient_df(indian_df) indian_ingredient_df.head(10).plot.barh()

韩国菜:

korean_ingredient_df = create_ingredient_df(korean_df) korean_ingredient_df.head(10).plot.barh()

对比五张图可以发现:rice(米饭)、garlic(蒜)、ginger(姜)几乎在每个菜系的头部都榜上有名——它们是跨菜系的公共特征,对"区分菜系"这个目标贡献很小,反而会稀释判别性特征的信号。

4.3 剔除混淆性公共特征

调用drop()移除最容易在菜系之间造成混淆的高频食材(原文档幽默地写道:Everyone loves rice, garlic and ginger!):

feature_df = df.drop(['cuisine','Unnamed: 0','rice','garlic','ginger'], axis=1) labels_df = df.cuisine #.unique() feature_df.head()

注意dropaxis=1表示按列删除。这一步是典型的**特征筛选(feature selection)**实践:删除在所有类中近似同分布的特征,让模型专注于真正有区分度的维度。剔除 4 列后,feature_df保留 381 个特征(1 个Unnamed: 0也已移除,实际进入建模的特征为 380 个食材列)。

五、实战步骤 4:用 SMOTE 均衡数据集

数据清理完成后,使用SMOTE(Synthetic Minority Over-sampling Technique,合成少数类过采样技术)来均衡类别分布。SMOTE 是 imbalanced-learn 库中imblearn.over_sampling模块提供的算法,其策略是通过插值生成新样本来增加少数类的数量,而不是简单复制已有样本。

5.1 执行过采样

调用fit_resample()

oversample = SMOTE() transformed_feature_df, transformed_label_df = oversample.fit_resample(feature_df, labels_df)

fit_resample()接收特征矩阵与标签序列两个参数,返回重采样后的特征矩阵与标签序列。均衡的意义在于:以二分类为例,若大部分数据属于某一个类,ML 模型会仅仅因为该类样本多而更频繁地预测它;均衡操作消除这种偏斜,让模型对每个类"一视同仁"。

5.2 对比均衡前后标签数量

print(f'new label count: {transformed_label_df.value_counts()}') print(f'old label count: {df.cuisine.value_counts()}')

输出:

new label count: korean 799 chinese 799 indian 799 japanese 799 thai 799 Name: cuisine, dtype: int64 old label count: korean 799 indian 598 chinese 442 japanese 320 thai 289 Name: cuisine, dtype: int64

SMOTE 以样本最多的类(korean,799 条)为基准,把其余四类全部补齐到 799 条。对仓库中实际产出的 cleaned_cuisines.csv 进行验证可以确认:该文件为 3995 行(5 × 799)、382 列(1 个 cuisine 标签列 + 381 个特征列),每个菜系恰好 799 条,与文档描述完全一致。

5.3 合并并导出均衡后的数据

最后一步把标签与特征合并为一个可导出的完整 DataFrame,并检查后保存:

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")

保存完成后,仓库的 data/cleaned_cuisines.csv 即为后续课程(2-Classifiers-1、3-Classifiers-2 等)直接加载的输入数据。至此,数据干净、均衡且"非常美味"。

六、进阶挑战、自修与课后作业

6.1 挑战(Challenge)

本课程包含多个有趣的数据集。翻阅各单元的data文件夹(如 2-Regression/data/US-pumpkins.csv、5-Clustering/data/nigerian-songs.csv),看看哪些适合做二分类或多分类问题,并写出你想向该数据提出的具体问题。

6.2 回顾与自修

研读 SMOTE 的 API。思考:它最适合什么使用场景?它解决的是什么问题?(提示:过采样合成样本的代价与适用边界,例如它只对数值特征有意义,且高维稀疏特征下需谨慎评估合成样本的有效性。)

6.3 课后作业

完成 探索分类方法:在 Scikit-learn 的监督学习文档中寻找分类算法,做一场"寻宝"——为课程中的某个数据集、一个可以提出来的问题、一种分类技术建立对应关系,整理成表格并解释该数据集如何与该分类算法配合使用。评分标准要求至少概述 3 种算法(优秀档要求 5 种),且解释要详细、准确。

6.4 多语言版本

本课时(课程编号第 10 课)还提供 R 语言版本,参考 solution/R/lesson_10.html,配套的 R 版清洗数据为 data/cleaned_cuisines_R.csv。

七、本课小结:数据准备四板斧

这一课表面是"分类入门",实质演示了经典机器学习中分类项目最前置、也最决定成败的数据准备流水线:

  1. 诊断head()/info()/value_counts()快速摸清数据规模(2448×385)、列类型与类别分布;
  2. 特征理解:用转置 + 求和的create_ingredient_df()提取每个类的 Top 特征,识别跨类公共特征;
  3. 特征清洗drop(['rice','garlic','ginger'])剔除混淆性高公共特征;
  4. 类别均衡SMOTE().fit_resample()将少数类插值补齐到多数类水平,产出 3995 行均衡数据集。

产出的 cleaned_cuisines.csv 是贯穿整个分类单元的数据基石,下一课将加载它,用 Logistic 回归、SVM 等分类器正式回答那个"多分类问题":给定一组食材,它属于哪种菜系。

【免费下载链接】ML-For-Beginners12 weeks, 26 lessons, 52 quizzes, classic Machine Learning for all项目地址: https://gitcode.com/GitHub_Trending/ml/ML-For-Beginners

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询