小白python入门 - 68. 分类入门
小白python入门 - 68. 分类入门
0. 写给同学的话
前两课解决了「地图」和「怎么干净地喂数据」。这一课进入监督学习里最常见的一类问题:分类——预测样本属于哪个类别。
我们用两个好懂的算法开胃:
- k 近邻(kNN):看最近的 k 个邻居怎么投票
- 逻辑回归(Logistic Regression):名字带「回归」,干的却是分类;核心是sigmoid把分数压成 0~1 的概率感
然后重点学怎么评价分类模型:混淆矩阵、精确率、召回率、F1,以及「准确率 99% 也可能是垃圾」的陷阱。
| 项目 | 说明 |
|---|---|
| 本节目标 | 理解 kNN 投票与逻辑回归 sigmoid;会读混淆矩阵与 P/R/F1;会对比两种模型的 sklearn 代码 |
| 学完能干什么 | 完成二分类小实验,并会用比 accuracy 更完整的指标汇报 |
| 预计时间 | 阅读 45–70 分钟;敲代码 40–50 分钟 |
| 前置 | 66 划分与 fit;67 Pipeline/缩放(kNN 强烈建议缩放) |
1. 生活里的例子 / 背景
1.1 分类问题长什么样?
| 问题 | 类别例子 | 备注 |
|---|---|---|
| 邮件是否垃圾 | 是 / 否 | 二分类 |
| 肿瘤良恶性 | 良性 / 恶性 | 二分类,错判代价不对称 |
| 手写数字 | 0–9 | 多分类 |
| 课程满意度 | 差 / 中 / 好 | 多分类(有时当有序) |
回归 vs 分类再钉一次:
- 问「是多少」→ 回归(下一课)
- 问「是哪一类」→ 分类(本课)
1.2 kNN 的生活版:近朱者赤
你搬进一个新宿舍楼,想猜「隔壁同学爱不爱打篮球」。你观察:住得最近的 5 个人里,4 个每周打球 → 你猜他也爱打。
这就是k 近邻:不先学一个复杂公式,而是临时查邻居。
1.3 逻辑回归的生活版:把「倾向分」压成概率
老师根据出勤、作业给一个「挂科风险分」:分数越高越危险。但业务要的是「挂科概率 0~1」,好设定阈值(比如 >0.5 预警)。
sigmoid曲线像一个软开关:把任意实数压到 (0,1) 之间。
1.4 为什么准确率会骗人?
1000 封邮件里只有 10 封垃圾。模型永远预测「正常」:
- 正确 990 封 → 准确率 99%
- 但垃圾一封没抓到 → 对反垃圾系统毫无用处
所以分类一定要会看混淆矩阵和精确率/召回率。
2. 核心概念(白话 + 表格)+ 相关图
2.1 k 近邻(k-Nearest Neighbors)
| 要点 | 大白话 |
|---|---|
| 思想 | 新样本的类别 ≈ 特征空间里最近的 k 个训练样本的多数类 |
| k | 看几个邻居;k 太小易噪声,k 太大易糊成「随大流」 |
| 距离 | 常用欧氏距离;特征量纲差大时必须先缩放(见 67 课) |
| 懒惰学习 | 训练几乎只是「记住数据」,计算主要在预测时 |
优点:好懂、能做非线性边界、超参少。
缺点:样本多时预测慢;高维距离失效;对无关特征和缩放敏感。
黑话:
- 决策边界:平面/空间里「判成 A 还是 B」的分界线
- 超参数:不能靠 fit 直接学、要你指定的数(如 k)
2.2 逻辑回归(其实是分类器)
名字历史原因带 Regression,任务是分类。
核心两步(二分类):
- 先算线性分:(z = w_1 x_1 + w_2 x_2 + \cdots + b)
- 再用sigmoid压到 (0,1):(\sigma(z) = \dfrac{1}{1+e^{-z}})
| z 很大正 | σ(z) 接近 1 | 模型更倾向正类 |
|---|---|---|
| z=0 | σ=0.5 | 中间地带 |
| z 很大负 | σ 接近 0 | 更倾向负类 |
默认常把概率 ≥ 0.5 判为正类(阈值可按业务改)。
| 对比 | kNN | 逻辑回归 |
|---|---|---|
| 在学什么 | 几乎存数据 + 距离投票 | 学一组权重 w 和偏置 b |
| 输出 | 类别(也可看邻居比例) | 自然带概率感 |
| 缩放 | 通常需要 | 建议做,便于收敛与解释 |
| 可解释 | 弱(「因为邻居这样」) | 相对强(特征权重方向) |
| 数据很大 | 预测可能慢 | 通常更快 |
2.3 混淆矩阵(Confusion Matrix)
以二分类、正类=「有病/是垃圾/会流失」为例:
| 预测:负 | 预测:正 | |
|---|---|---|
| 真实:负 | TN 真负 | FP 假正(误报) |
| 真实:正 | FN 假负(漏报) | TP 真正 |
| 符号 | 英文 | 大白话 |
|---|---|---|
| TP | True Positive | 真有问题,也抓对了 |
| TN | True Negative | 真没事,也放对了 |
| FP | False Positive | 没事却报警(狼来了) |
| FN | False Negative | 有事却漏了(漏诊) |
2.4 精确率、召回率、F1、准确率
用 TP/FP/FN 定义(先抓直觉,公式为辅):
| 指标 | 公式直觉 | 大白话 | 在意谁 |
|---|---|---|---|
| 准确率 Accuracy | 全体对的比例 | 整体蒙对多少 | 类别均衡时还行 |
| 精确率 Precision | TP / (TP+FP) | 你报「正」的里面有多少真是正 | 讨厌误报时 |
| 召回率 Recall | TP / (TP+FN) | 所有真正的正里抓回多少 | 讨厌漏报时 |
| F1 | 精确率与召回的调和平均 | 两者平衡的一个分数 | 综合看 |
业务口诀:
- 垃圾邮件:有时宁可多进垃圾箱(召回高)或宁可少误杀(精确高)——产品定
- 癌症筛查:通常更怕FN 漏诊→ 重视召回
- 广告投放「高意向」:更怕FP 浪费预算→ 重视精确
2.5 准确率陷阱(再强调)
| 场景 | 瞎猜策略 | Accuracy | 是否有用 |
|---|---|---|---|
| 99% 负类 | 全猜负 | ~99% | 常没用 |
| 均衡 50/50 | 全猜一类 | ~50% | 基线参考 |
汇报建议:至少同时给混淆矩阵 + precision/recall/F1;类别不均衡时优先看后者,或使用class_weight、重采样等(后文进阶)。
3. 算法 / 流程用图说清楚
3.1 kNN 预测一步步
1. 准备:训练集已缩放(重要) 2. 来一个新点 x 3. 算 x 到所有训练点的距离 4. 取最近的 k 个点 5. 看这 k 个点的标签:多数表决 → 预测类别 (也可看各类占比当「软」结果)3.2 逻辑回归训练与预测(直觉版)
训练: 反复调整 w, b 让「预测概率」更贴近真实标签 (内部用优化算法,本课不展开公式推导) 预测: z = w·x + b p = sigmoid(z) 若 p >= 阈值(默认 0.5)→ 正类,否则负类3.3 本课实验总流程
加载二分类数据(乳腺癌) → train_test_split(分层 stratify) → Pipeline(StandardScaler + 模型) → fit 训练集 → 测试集:accuracy / 混淆矩阵 / classification_report → 对比 kNN vs 逻辑回归4. 手把手环境与代码
4.1 安装
pipinstall-Uscikit-learn pandas numpy4.2 数据:乳腺癌二分类(sklearn 自带)
特征是细胞核相关测量值,标签:恶性/良性。仅作教学,不能当真实医疗结论。
fromsklearn.datasetsimportload_breast_cancerfromsklearn.model_selectionimporttrain_test_splitfromsklearn.preprocessingimportStandardScalerfromsklearn.pipelineimportPipelinefromsklearn.neighborsimportKNeighborsClassifierfromsklearn.linear_modelimportLogisticRegressionfromsklearn.metricsimport(accuracy_score,confusion_matrix,classification_report,precision_score,recall_score,f1_score,)data=load_breast_cancer()X,y=data.data,data.targetprint("特征维度:",X.shape)print("类别:",list(zip([0,1],data.target_names)))print("各类数量:",{data.target_names[i]:int((y==i).sum())foriin[0,1]})X_train,X_test,y_train,y_test=train_test_split(X,y,test_size=0.25,random_state=42,stratify=y)预期输出形态:形状约(569, 30);两类名称与计数(恶性/良性数量不完全相等)。
stratify=y:分层抽样,让训练/测试里正负比例接近总体,避免「测试集碰巧几乎全是一类」。
4.3 kNN 完整评估
pipe_knn=Pipeline([("scaler",StandardScaler()),("clf",KNeighborsClassifier(n_neighbors=5)),])pipe_knn.fit(X_train,y_train)pred_knn=pipe_knn.predict(X_test)print("=== kNN ===")print("Accuracy:",round(accuracy_score(y_test,pred_knn),4))print("Confusion matrix:\n",confusion_matrix(y_test,pred_knn))print(classification_report(y_test,pred_knn,target_names=data.target_names))预期:Accuracy 通常较高(玩具数据);混淆矩阵 2×2;report 里有 precision/recall/f1。
读矩阵:confusion_matrix默认行=真实、列=预测(与 sklearn 文档一致)。先确认再解读 TP/FP。
4.4 逻辑回归完整评估 + 概率
pipe_lr=Pipeline([("scaler",StandardScaler()),("clf",LogisticRegression(max_iter=5000)),])pipe_lr.fit(X_train,y_train)pred_lr=pipe_lr.predict(X_test)proba_lr=pipe_lr.predict_proba(X_test)[:,1]# 正类概率print("=== Logistic Regression ===")print("Accuracy:",round(accuracy_score(y_test,pred_lr),4))print("Confusion matrix:\n",confusion_matrix(y_test,pred_lr))print(classification_report(y_test,pred_lr,target_names=data.target_names))print("前 5 个样本的正类概率:",[round(p,3)forpinproba_lr[:5]])print("前 5 个预测标签:",pred_lr[:5])print("前 5 个真实标签:",y_test[:5])预期:概率是 0~1 小数;阈值 0.5 时,概率高的对应预测 1。
4.5 并排对比(同一划分、同一套指标)
defsummarize(name,y_true,y_pred,pos_label=1):print(f"\n{name}")print(" acc =",round(accuracy_score(y_true,y_pred),4))print(" precision =",round(precision_score(y_true,y_pred,pos_label=pos_label),4))print(" recall =",round(recall_score(y_true,y_pred,pos_label=pos_label),4))print(" f1 =",round(f1_score(y_true,y_pred,pos_label=pos_label),4))print(" cm =\n",confusion_matrix(y_true,y_pred))summarize("kNN k=5",y_test,pred_knn)summarize("LogReg",y_test,pred_lr)怎么读结果:不必纠结谁高 0.01;关注流程是否正确、指标是否全面。换random_state或 k,名次可能对调。
4.6 改 k 看趋势(小实验)
forkin[1,3,5,15,33]:pipe=Pipeline([("scaler",StandardScaler()),("clf",KNeighborsClassifier(n_neighbors=k)),])pipe.fit(X_train,y_train)acc=accuracy_score(y_test,pipe.predict(X_test))print(f"k={k:2d}test_acc={acc:.4f}")预期趋势直觉:k=1 可能波动大;k 过大可能变钝。具体数字自己跑。
4.7 准确率陷阱:人造不均衡数据
importnumpyasnpfromsklearn.dummyimportDummyClassifier rng=np.random.RandomState(0)n=1000# 95% 为类别 0y_imbal=(rng.rand(n)>0.95).astype(int)X_imbal=rng.randn(n,5)Xtr,Xte,ytr,yte=train_test_split(X_imbal,y_imbal,test_size=0.3,random_state=0,stratify=y_imbal)dummy=DummyClassifier(strategy="most_frequent")dummy.fit(Xtr,ytr)pred_d=dummy.predict(Xte)print("多数类瞎猜 Accuracy:",round(accuracy_score(yte,pred_d),4))print("混淆矩阵:\n",confusion_matrix(yte,pred_d))print(classification_report(yte,pred_d,zero_division=0))预期:Accuracy 可以很高;但少数类 recall 经常是 0。这就是「准确率陷阱」的数字版。
4.8 阈值不是只能 0.5(开拓视野)
# 以逻辑回归概率为例:把阈值改成 0.3,更易判成正类 → 召回往往升、精确往往降thr=0.3pred_thr=(proba_lr>=thr).astype(int)print(f"阈值={thr}")print(classification_report(y_test,pred_thr,target_names=data.target_names))业务若更怕漏报,可降低正类阈值(在验证集上选,不要只在测试集上反复抠——71 课再系统讲)。
5. 常见坑(大白话)
| 坑 | 现象 | 正确直觉 |
|---|---|---|
| kNN 不缩放 | 距离被大数值特征绑架 | Pipeline 加StandardScaler |
| 只报 Accuracy | 不均衡时自我感觉良好 | 看 cm + P/R/F1 |
| 搞反 precision/recall | 和业务对着干 | 精确=报得准;召回=抓得全 |
| 混淆矩阵行列搞反 | 故事讲反 | 先查文档:行真实、列预测 |
逻辑回归没max_iter够 | 收敛警告 | 加大max_iter或先缩放 |
| 把逻辑回归当「因果解释」 | 权重当因果 | 相关预测 ≠ 因果 |
| 测试集上疯狂调 k/阈值 | 测试分泄漏式虚高 | 验证集或交叉验证(71) |
| 多分类仍用二分类口径乱讲 | 指标对不上 | 多分类看 macro/weighted 等平均方式 |
| 医疗/金融直接上线玩具模型 | 伦理与合规风险 | 教学数据 ≠ 生产决策 |
6. 对照表、小结
6.1 算法速查
| kNN | 逻辑回归 | |
|---|---|---|
| 核心 | 邻居投票 | 线性分 + sigmoid |
| 关键超参 | k、距离 | 正则强度 C 等(后文) |
| 概率 | 间接 | predict_proba自然 |
| 缩放 | 重要 | 建议 |
6.2 指标速查
| 指标 | 一句话 |
|---|---|
| Accuracy | 总体对了多少 |
| Precision | 报正里有多少真对 |
| Recall | 真正当中抓回多少 |
| F1 | 精确与召回的平衡 |
| 混淆矩阵 | TP/FP/FN/TN 一张表看清 |
6.3 本节三句话
- 分类猜「哪一类」;kNN 靠邻居,逻辑回归靠加权分 + sigmoid。
- 类别不均衡时,准确率可以很好看却没用。
- 汇报请带上混淆矩阵和 precision/recall/F1,并对齐业务更怕误报还是漏报。
7. 小练笔(由易到难)
题 1
用食堂例子解释 k=1 和 k=9 可能有什么不同(邻居太少 vs 太多)。
题 2
sigmoid 输入从 -10 变到 +10,输出大概从什么范围变到什么范围?是否可能等于 0 或 1?
题 3
某混淆矩阵:TP=40, FP=10, FN=20, TN=930。
手算 Accuracy、Precision、Recall(正类)。体会准确率为何仍可能「看起来不错」。
题 4(代码)
在乳腺癌数据上,比较n_neighbors=1与n_neighbors=25的测试 F1,并打印两张混淆矩阵。
题 5
反垃圾邮件系统:领导说「不能漏掉垃圾邮件」,产品说「千万别把重要邮件丢进垃圾箱」。
更该优先盯 precision 还是 recall?两者冲突时你怎么跟领导用混淆矩阵沟通?
参考思路:
1 k=1 跟最近点走,噪点敏感;k=9 更平滑但可能模糊。2 从接近 0 到接近 1,一般达不到绝对 0/1。3 Acc=(40+930)/1000=0.97;P=40/50=0.8;R=40/60≈0.67。4 自跑。5 不漏 → 重视召回;不误杀 → 重视精确;用 FP/FN 代价谈阈值。
8. 下一课预告
69. 回归入门
分类问「是哪一类」,回归问「是多少」。下一课用线性回归预测连续值,认识 MAE/RMSE/R²,并初步碰到多项式过拟合与 Ridge/Lasso 正则化——和 66 课的过拟合地图对上号。
引用与参考
- scikit-learn 监督学习总览:https://scikit-learn.org/stable/supervised_learning.html
KNeighborsClassifier:https://scikit-learn.org/stable/modules/generated/sklearn.neighbors.KNeighborsClassifier.htmlLogisticRegression:https://scikit-learn.org/stable/modules/generated/sklearn.linear_model.LogisticRegression.html- 分类指标:https://scikit-learn.org/stable/modules/model_evaluation.html#classification-metrics
- 混淆矩阵:https://scikit-learn.org/stable/modules/generated/sklearn.metrics.confusion_matrix.html
- Breast cancer 数据集:https://scikit-learn.org/stable/modules/generated/sklearn.datasets.load_breast_cancer.html
- Wikipedia - k-nearest neighbors algorithm:https://en.wikipedia.org/wiki/K-nearest_neighbors_algorithm
- Wikipedia - Logistic regression:https://en.wikipedia.org/wiki/Logistic_regression
- Google MLCC - 分类:https://developers.google.com/machine-learning/crash-course/classification
再讲一个完整小故事(指标怎么选)
学校心理中心做一个粗筛:问卷分数预测「是否需要人工回访」。
- 若召回率低:真正需要帮助的人被漏掉 → 风险大。
- 若精确率低:大量健康同学被误报 → 人工忙不过来,同学也被吓到。
这时你不会只说「准确率 95% 真棒」,而会问:
- 漏掉一个,代价是什么?
- 误报一个,代价是什么?
- 默认阈值 0.5 要不要调低一点,宁可多回访?
分类课的毕业标准:你会看混淆矩阵,会用大白话解释精确率/召回率,会在业务场景里做取舍——而不是只会打印 accuracy。
课堂讨论题(可分组 10 分钟)
- 如果老板只要一个数字「准确率」,你怎么用两分钟说服他看混淆矩阵?
- 数据只有 80 条,你还上机器学习吗?为什么?
- 你更愿意维护:100 条清晰业务规则,还是一个 90 分但没人能解释的模型?
把讨论结论写在笔记里——比多抄 50 行 API 更接近真实工作。
自我检测(不看稿,口头答)
- 我能不看笔记讲清本课最重要的一张图在说什么
- 我能指出一段「错误代码」错在哪
- 我能举一个生活例子对应本课任务类型
- 我知道下一课大概要解决什么痛点
全部打勾再进入下一课,效率更高。
附录:给大一的 FAQ(本课补充)
下面这些问题,是第一次学本课内容时最容易卡住的地方。用白话再过一遍。
Q1:我是不是一定要背公式?
不必先背公式。你要先会讲故事:输入是什么、输出是什么、模型在怕什么(过拟合、泄漏、指标骗人)。公式是为了精确表达故事;故事通了,公式只是翻译。
Q2:代码跑不通怎么办?
按这个顺序排查:
- 虚拟环境激活了吗?(提示符前有没有 .venv)
- 包装了吗?python -c “import sklearn; print(sklearn.version)”
- 报错最后一行是什么?把Error 类型 + 最后一行记下来再搜
- 路径、文件名、中文引号有没有混用
- 仍不行:换一个最小例子(本课最前面的 10 行代码)确认环境 OK
Q3:我和同学分数差很多,是不是我很差?
不一定。可能是:
andom_state 不同、数据划分不同、指标不同、甚至泄漏导致虚高。先对齐评估协议,再比分数。
Q4:这课和「人工智能 / ChatGPT」是什么关系?
ChatGPT 一类是很大的深度学习系统,偏语言与对话。本课练的是表格/经典机器学习基本功:分类、回归、评估、Pipeline。基本功会了,你以后学深度学习或用大模型 API,才知道自己在解决什么问题、如何公平比较。
Q5:我需要买 GPU 吗?
本阶段不需要。sklearn 在普通笔记本 CPU 上就够。GPU 主要是深度学习训练时才刚需。
Q6:作业要做到什么程度算合格?
最低标准:
- 能用自己的话讲清本课 3 个核心概念
- 能跑通本课主线代码,并看懂输出含义
- 能指出至少 2 个常见坑
- 小练笔完成一半以上(鼓励全做)
Q7:我想继续深入,课外看什么?
优先官方文档对应章节(见文末引用),其次 ISLR 中文/英文入门章节。别一上来就啃很厚的证明书——容易劝退。
本课概念速记卡(可抄笔记本)
| 我用大白话怎么说 | 对应术语 |
|---|---|
| 用历史数据猜新情况 | 机器学习 / 预测 |
| 拿来学的那部分数据 | 训练集 |
| 假装是新客户的那部分 | 测试集 |
| 背答案背过头 | 过拟合 |
| 笨到学不会 | 欠拟合 |
| 偷看了考题 | 数据泄漏 |
| 步骤焊成一条龙 | Pipeline |
建议学习节奏
| 时间 | 做什么 |
|---|---|
| 第 1 小时 | 只读例子与图,不写代码 |
| 第 2 小时 | 抄跑主线代码,改一个参数观察变化 |
| 第 3 小时 | 做小练笔 + 写 5 句笔记 |
| 之后 | 隔一天不看稿子复述一遍 |
记住:大一阶段,「讲清楚 + 跑得通 + 知道坑」比「一次记住全部 API」重要得多。
