引言
分类是监督学习的另一大类:输出不再是连续数值,而是离散类别——垃圾邮件(垃圾/正常)、欺诈交易(欺诈/正常)、肿瘤(良性/恶性)。分类问题的难点不在「训练一个模型」,而在评估与权衡:假阳性与假阴性哪个更不能接受?数据严重不平衡时准确率会骗人,怎么办?
本文用信用卡欺诈检测这一经典不平衡场景,从逻辑回归讲起,逐个跑通 KNN、决策树、随机森林,重点落在分类评估体系——混淆矩阵、精确率/召回率/F1、AUC——并用这些指标反哺模型选择。最后给出类别不平衡的处理套路。
前置:[[ml]] 专题的环境搭建与回归基础(https://plumephp.com/ml-supervised-regression/)。深度算法原理见 [[ai-ml]] 专题。
目录
- 1. 分类问题与逻辑回归的直觉
- 2. 数据准备:信用卡欺诈检测
- 3. 第一个分类模型:逻辑回归
- 4. 混淆矩阵:看清对错分布
- 5. 评估指标:精确率、召回率、F1 与 AUC
- 6. 更多分类器:KNN、决策树与随机森林
- 7. 类别不平衡的处理
- 8. 阈值调优与业务权衡
- 9. 总结:分类问题的完整套路
- 延伸阅读
1. 分类问题与逻辑回归的直觉
1.1 从回归到分类
线性回归输出连续值;分类需要输出「属于某类的概率」。逻辑回归在线性组合上套一个 Sigmoid 函数,把任意实数值压缩到 (0,1),作为概率:
p = 1 / (1 + exp(-(w·x + b)))
p > 0.5 → 预测为正类
p < 0.5 → 预测为负类
1.2 决策边界
逻辑回归学到的是一条「分界线」,把特征空间切成正/负两侧。虽然名字带「回归」,但它是分类器。
2. 数据准备:信用卡欺诈检测
2.1 数据集概览
我们用 UCI 信用卡欺诈数据(或构造等价结构):28 个匿名特征 + Class(0=正常,1=欺诈)。
import pandas as pd
from sklearn.model_selection import train_test_split
df = pd.read_csv('creditcard.csv')
print(df.shape) # 通常约 28 万行
print(df['Class'].value_counts(normalize=True))
# 0 (正常): 99.8%
# 1 (欺诈): 0.2% ← 极度不平衡
2.2 特征与目标切分
X = df.drop(columns=['Class'])
y = df['Class']
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.2, random_state=42, stratify=y)
# stratify=y:切分后训练/测试集类别比例一致
3. 第一个分类模型:逻辑回归
3.1 训练
from sklearn.linear_model import LogisticRegression
model = LogisticRegression(max_iter=2000)
model.fit(X_train, y_train)
3.2 为什么直接看准确率会「被骗」
from sklearn.metrics import accuracy_score
y_pred = model.predict(X_test)
print("准确率:", accuracy_score(y_test, y_pred).round(4))
# 很可能 99.8%+ —— 因为 99.8% 都是正常交易
# 「全预测为正常」也能拿到 99.8% 准确率,但这个模型毫无价值
不平衡数据下,准确率不是好指标。必须看每一类的表现,尤其是少数类。
4. 混淆矩阵:看清对错分布
4.1 四象限
预测\真实 正(欺诈) 负(正常)
预测正 TP(真阳) FP(假阳)
预测负 FN(假阴) TN(真阴)
- TP:真的欺诈被抓住
- FP:正常被误判为欺诈(打扰用户)
- FN:真欺诈漏掉了(更严重)
from sklearn.metrics import confusion_matrix
cm = confusion_matrix(y_test, y_pred)
print(cm)
4.2 用 pandas 美化输出
import seaborn as sns
import matplotlib.pyplot as plt
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues')
plt.xlabel('预测')
plt.ylabel('真实')
plt.show()
5. 评估指标:精确率、召回率、F1 与 AUC
5.1 三大核心指标
| 指标 | 公式 | 直觉 | 关注点 |
|---|---|---|---|
| 精确率 Precision | TP/(TP+FP) | 预测为正的里有多少真对 | 别误伤正常 |
| 召回率 Recall | TP/(TP+FN) | 真实正类里抓到多少 | 别漏掉欺诈 |
| F1 | 2·P·R/(P+R) | 精确率与召回率的调和平均 | 综合平衡 |
from sklearn.metrics import precision_score, recall_score, f1_score
print("精确率:", precision_score(y_test, y_pred).round(4))
print("召回率:", recall_score(y_test, y_pred).round(4))
print("F1 :", f1_score(y_test, y_pred).round(4))
5.2 该看哪个:取决于业务
| 业务场景 | 看重指标 | 原因 |
|---|---|---|
| 欺诈检测 | 召回率 | 漏掉欺诈损失巨大 |
| 垃圾邮件 | 精确率 | 误删正常邮件不可接受 |
| 通用 | F1 | 平衡 |
5.3 AUC:排序能力的全局指标
roc_auc_score 衡量「把正类排在负类前面的能力」,取值 0.5(随机)~1(完美),与阈值无关:
from sklearn.metrics import roc_auc_score
y_proba = model.predict_proba(X_test)[:, 1]
print("AUC:", roc_auc_score(y_test, y_proba).round(4))
6. 更多分类器:KNN、决策树与随机森林
6.1 数据量大先标准化
KNN 依赖距离,特征量级不一致会失真;随机森林不依赖。
from sklearn.preprocessing import StandardScaler
scaler = StandardScaler()
X_train_s = scaler.fit_transform(X_train)
X_test_s = scaler.transform(X_test)
6.2 KNN:懒惰学习
from sklearn.neighbors import KNeighborsClassifier
knn = KNeighborsClassifier(n_neighbors=5)
knn.fit(X_train_s, y_train)
print("KNN F1:", f1_score(y_test, knn.predict(X_test_s)).round(4))
6.3 决策树:可解释的分支规则
from sklearn.tree import DecisionTreeClassifier
tree = DecisionTreeClassifier(max_depth=5, random_state=42)
tree.fit(X_train, y_train)
# 可以导出树结构做解释(业务可解释性)
6.4 随机森林:多棵树的投票
from sklearn.ensemble import RandomForestClassifier
rf = RandomForestClassifier(n_estimators=100, random_state=42)
rf.fit(X_train, y_train)
print("RF F1:", f1_score(y_test, rf.predict(X_test)).round(4))
6.5 特征重要性
importance = pd.Series(rf.feature_importances_, index=X.columns)
print(importance.sort_values(ascending=False).head(10))
6.6 模型对比
| 模型 | 优点 | 缺点 |
|---|---|---|
| 逻辑回归 | 简单、可解释、快 | 线性假设 |
| KNN | 无需训练 | 慢、维度灾难 |
| 决策树 | 可解释 | 易过拟合 |
| 随机森林 | 强、抗过拟合 | 难解释、慢 |
7. 类别不平衡的处理
7.1 三种常用手段
# 方法1:模型内置权重(class_weight 惩罚少数类误分)
LogisticRegression(class_weight='balanced')
# 方法2:欠采样/过采样(简单)
from imblearn.over_sampling import RandomOverSampler
# 或更优的 SMOTE 合成少数类
# 方法3:阈值调优(见下一节)
7.2 SMOTE 合成样本
from imblearn.over_sampling import SMOTE
smote = SMOTE(random_state=42)
X_res, y_res = smote.fit_resample(X_train, y_train)
print(pd.Series(y_res).value_counts(normalize=True)) # 已均衡
7.3 处理后重新评估
model.fit(X_res, y_res)
print("SMOTE后 召回率:", recall_score(y_test, model.predict(X_test)).round(4))
注意:只对训练集做重采样,测试集保持真实分布,否则评估失真。
8. 阈值调优与业务权衡
8.1 默认阈值 0.5 不一定最优
把阈值调低,会抓到更多欺诈(召回↑),但也误伤更多正常(精确↓)。用 ROC 曲线可视化这个权衡:
from sklearn.metrics import roc_curve
fpr, tpr, thresholds = roc_curve(y_test, y_proba)
plt.plot(fpr, tpr)
plt.xlabel('假阳率 (FPR)')
plt.ylabel('真阳率 (TPR)')
plt.title('ROC 曲线')
plt.show()
8.2 按业务目标选阈值
# 找「召回率 ≥ 0.9」的最低阈值
import numpy as np
idx = np.where(tpr >= 0.9)[0][0]
threshold = thresholds[idx]
print("阈值:", round(threshold, 4), "对应召回率:", round(tpr[idx], 4))
9. 总结:分类问题的完整套路
9.1 九步模板
# 1. 切分(stratify 保持类别比例)
# 2. 标准化(距离型模型必需)
# 3. 训练基线模型(逻辑回归)
# 4. 看混淆矩阵,别只看准确率
# 5. 按业务选指标(精确率/召回率/F1/AUC)
# 6. 上更强模型(随机森林)
# 7. 处理不平衡(class_weight / SMOTE)
# 8. 阈值调优(ROC)
# 9. 交叉验证定参,最终评估测试集
9.2 关键决策点
| 问题 | 选择 |
|---|---|
| 数据不平衡? | 别用准确率;看召回/F1 |
| 业务重漏杀 | 高召回率 + 调低阈值 |
| 业务重误伤 | 高精确率 + 调高阈值 |
| 要可解释 | 逻辑回归/决策树 |
| 要性能 | 随机森林/XGBoost |
9.3 一句话心法
分类的成败不在模型多炫,而在「用对指标、平衡好数据、调好阈值」。先跑通逻辑回归基线,再逐步升级。
延伸阅读
- https://plumephp.com/ml-supervised-regression/ — 回归基础(同一套训练/评估思维)
- https://plumephp.com/ml-model-evaluation/ — 更系统的模型评估与选择
- [[ai-ml]] 专题的监督学习与模型评估深度文章
- scikit-learn 分类指标文档
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。