首页
/ ML-For-Beginners 分类入门:基于亚洲与印度料理数据集的多分类数据清洗与 SMOTE 平衡实战

ML-For-Beginners 分类入门:基于亚洲与印度料理数据集的多分类数据清洗与 SMOTE 平衡实战

2026-09-06 15:32:12作者:尤峻淳Whitney

本文对应 ML-For-Beginners 课程《Getting started with classification》模块的第一课(课程第 10 课,原文档 README.md,该文档为英文原版 4-Classification/1-Introduction/README.md 的阿拉伯语机器翻译版)。本篇将带你把"亚洲与印度料理"数据集(2448 道菜谱、385 个成分特征列)加工为一个特征去噪、类别均衡的多分类训练集:从理解二分类与多分类的本质区别,到用 Pandas 探查数据分布,再到用 imblearn 的 SMOTE 过采样把 289~799 条不均衡样本拉平到每类 799 条,并导出 cleaned_cuisines.csv 供后续三课的分类算法使用。

1. 分类的定位:监督学习的另一面

分类(Classification)是经典机器学习中与回归并列的核心任务,属于监督学习(Supervised Learning):数据带有标签,模型通过学习"输入特征 → 输出类别"的映射关系来建立预测模型。原文档特别强调它与回归的连续性,帮助读者建立知识衔接:

  • 线性回归(Linear Regression):预测变量间的连续关系,例如"同一款南瓜 9 月与 12 月的价格差异";
  • 逻辑回归(Logistic Regression):发现"二元类别",例如"在这个价格点,这个南瓜是橙色还是非橙色";
  • 分类:用多种算法确定数据点的标签或类别,并进一步分为 二分类(binary classification)多分类(multiclass classification) 两大族。
二分类与多分类问题对比
图 1:分类算法需要处理的二分类与多分类问题对比(信息图作者 Jen Looper)

原文档给出的思考题值得在动手前完成:想象一个料理数据集——多分类模型能回答什么问题?二分类模型又能回答什么问题?例如"判断某道菜是否很可能使用葫芦巴(fenugreek)"是二分类问题;而"给定一袋八角、朝鲜蓟、花椰菜和辣根,能否做出一道典型印度菜"则涉及多类别判断。

Scikit-learn 提供了多种分类算法,取决于要解决的问题类型。本模块共 4 课(见 4-Classification/README.md):

  1. Introduction to classification —— 本课:数据清洗与均衡;
  2. More classifiers
  3. Yet other classifiers
  4. Applied ML: build a web app

本课对应的课后作业是 assignment.md,要求查阅 Scikit-learn 文档中的分类方法清单,将算法与本课程数据集匹配并说明"将向数据提出什么问题"。

2. 任务定义:一个多分类问题

本课要回答的问题是:给定一组成分(ingredients),判断它属于哪种"菜系"。由于存在多个候选菜系(thai、japanese、chinese、indian、korean),这实际上是一个多分类问题

数据源为 4-Classification/data/cuisines.csv,其结构特点(可通过 solution/notebook.ipynb 与 CSV 表头确认):

  • 2448 行菜谱样本(不含表头),385 列
  • cuisine(字符串标签列,object 类型)外,384 列全部是 0/1 二值成分列(almondangelicaanisericesoy_sauce……),即每一列代表一种成分是否出现在该菜谱中;
  • 另有一列 Unnamed: 0(原始索引残留,需要丢弃)。

3. 环境准备与数据导入

动手前需要安装 imblearn(imbalanced-learn)。它是 Scikit-learn 生态中专门处理类别不均衡问题的包,本教程用它提供的 SMOTE 算法实现过采样。

pip install imblearn

然后在 notebook.ipynb(该目录下的空白 Notebook,位于本课根目录)中导入所需依赖:

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

导入数据集。注意:由于 Notebook 位于 4-Classification/1-Introduction/,数据目录在上一级的 data/ 文件夹,因此路径为 ../data/cuisines.csv(若你运行的是 solution 目录下的 solution/notebook.ipynb,其内部路径为 ../../data/cuisines.csv,同理):

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

read_csv() 会读取 cuisines.csv 的全部内容并放入变量 df。先用 head() 查看前五行:

df.head()

前五行形如(每行是一个菜谱,绝大多数成分列为 0,个别为 1):

|     | Unnamed: 0 | cuisine | almond | angelica | anise | anise_seed | apple | apple_brandy | ... | whiskey | white_bread | white_wine | ... | yogurt | zucchini |
| --- | ---------- | ------- | ------ | -------- | ----- | ---------- | ----- | ------------ | --- | ------- | ----------- | ---------- | --- | ------ | -------- |
| 0   | 65         | indian  | 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        |
| 2   | 67         | indian  | 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        |
| 4   | 69         | indian  | 0      | 0        | 0     | 0          | 0     | 0            | ... | 0       | 0           | 0          | ... | 1      | 0        |

再用 info() 获取整体元信息:

df.info()
<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 文件的实际规模:2448 条记录、385 列,其中 384 列为 int64 二值列、1 列为 object 标签列

4. 探查菜系分布:发现类别不均衡

先弄清数据在每个菜系上的分布情况:

df.cuisine.value_counts().plot.barh()
菜系数据分布
图 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)

直接对 cuisines.csvcuisine 列统计可复现该结果:korean 799、indian 598、chinese 442、japanese 320、thai 289,合计 2448。最大类(korean)是最小类(thai)的约 2.8 倍——这种偏斜会让模型"偏向多数类",正是第 6 节要用 SMOTE 修复的问题。

5. 成分画像:找出混淆各菜系的公共特征

深入数据之前,先看每种菜系的"典型成分"长什么样。原文档给出一个辅助函数 create_ingredient_df():先转置并丢弃 cuisineUnnamed: 0 两列,对每一行(即每个成分)求和得到该菜系内出现次数,再剔除全 0 成分,最后按出现次数降序排列:

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

solution/notebook.ipynb 中同一函数带有完整注释:df.T 转置后 sum(axis=1) 即每个成分在该菜系中的总出现次数,(ingredient_df.T != 0).any() 过滤掉零出现行。)

对每个菜系分别调用并绘制前 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)在几乎所有菜系中都是高频成分。从源码结构看,这些特征对区分"是哪国菜"几乎没有信息量,反而会成为跨菜系的噪声。因此用 drop() 将它们连同两个无信息列一并删除:

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

至此,feature_df 是 2448 × 380 的纯特征矩阵,labels_df 是 2448 个菜系标签。

6. 用 SMOTE 均衡数据集

数据已经"干净"了,但类别仍不均衡。这里引入 imblearnSMOTE(Synthetic Minority Over-sampling Technique,合成少数类过采样技术):它不是简单复制少数类样本,而是在特征空间中对少数类样本做插值,生成合成样本,从而把每类样本数提升到最大类的水平。

为什么必须均衡?以二分类为例:如果绝大多数数据属于某一种类,模型会仅因"该类的训练样本多"而更频繁地预测该类。SMOTE 消除这种偏斜,让分类结果反映真实判别能力。

调用 fit_resample() 执行过采样:

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

fit_resample() 返回两个对象:过采样后的特征矩阵与对应的标签序列。注意 SMOTE 要求特征为数值型——本数据集 380 列均为 0/1 二值特征,天然满足该前提。

对比新旧标签分布:

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

每个菜系都被抬升到 799 条,总量从 2448 变为 3995 条(799 × 5)。

最后一步:把均衡后的标签与特征合并为新的 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")

导出产物即仓库中已存在的 4-Classification/data/cleaned_cuisines.csv——直接检查该文件可以验证:共 3995 行数据,且 chinese、indian、japanese、korean、thai 各恰好 799 行,与上文 SMOTE 输出完全一致。后续第 11、12 课(2-Classifiers-13-Classifiers-2)都将基于这份均衡数据训练 Naive Bayes、SVM、Random Forest 等分类器。

7. 进阶练习与自研方向

原文档在收尾处给出了两个延伸任务:

  • 挑战题:翻查课程其他 data 文件夹(如 2-Regression/data/US-pumpkins.csv5-Clustering/data/nigerian-songs.csv),判断哪些数据集适合二分类或多分类,并列出你会向数据提出的问题;
  • 自研:探索 SMOTE 的 API——它的 k_neighborssampling_strategy 等参数在哪些场景下需要调整?它解决什么问题、又会在什么场景(例如高维稀疏特征)下带来风险?

另外,本课同时提供 R 语言版本(solution/R/lesson_10.html,含 lesson_10-R.ipynb),以及 Julia 说明(solution/Julia/README.md),可用其他语言对照同一工作流。

8. 小结

本课完成了分类项目的前置工程:

  1. 概念层:分类是监督学习的两大分支之一,按类别数分为二分类与多分类;本数据集是一个 5 类多分类问题;
  2. 探查层read_csvhead()/info() 确认 2448 × 385 的二值特征结构;value_counts().plot.barh() 暴露出 289~799 条的类别偏斜;
  3. 去噪层create_ingredient_df() 定位出 rice/garlic/ginger 等跨菜系高频成分并删除;
  4. 均衡层:SMOTE 过采样把每类拉到 799 条,导出 cleaned_cuisines.csv(3995 行)作为后续三课分类实验的统一输入。

这套"分布探查 → 公共特征剔除 → 合成过采样 → 落盘"的流水线,是所有类别不均衡分类项目的通用起手式。

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