已实现朴素贝叶斯,请教基于FMNIST绘制二/多分类ROC曲线(无scikit)
嘿,很高兴你已经搞定了朴素贝叶斯的核心部分!ROC曲线的关键其实是理解它本质是遍历所有可能的分类阈值,计算每个阈值下的真阳性率(TPR)和假阳性率(FPR),然后把这些点连起来。我一步步给你讲清楚,从二分类到多分类,再给你写可直接用的代码片段。
一、二分类场景(裤子 vs 套头衫)
核心思路
ROC曲线的前提是,你能给每个测试样本输出属于正类的概率值(而不是直接输出分类结果)。对你的朴素贝叶斯来说,就是对每个测试样本,计算它属于“裤子”(假设设为正类)的后验概率p_trouser,以及属于“套头衫”的后验概率p_pullover。我们只需要p_trouser来做ROC分析。
步骤1:准备数据
假设你已经有:
test_labels:测试集的真实标签(比如用字符串或数字区分裤子和套头衫)posteriors_test:测试集每个样本的后验概率数组,形状是(n_samples, 2),第一列是裤子的后验概率,第二列是套头衫的
步骤2:写辅助函数计算TPR和FPR
先封装一个函数,用来根据给定阈值计算对应的真阳性率和假阳性率:
import numpy as np def calculate_tpr_fpr(y_true, y_score, threshold): # y_true: 真实标签(正类标记为1,负类标记为0) # y_score: 每个样本属于正类的概率 # threshold: 当前判断阈值 y_pred = (y_score >= threshold).astype(int) # 计算混淆矩阵的四个指标 TP = sum((y_true == 1) & (y_pred == 1)) FP = sum((y_true == 0) & (y_pred == 1)) TN = sum((y_true == 0) & (y_pred == 0)) FN = sum((y_true == 1) & (y_pred == 0)) # 避免除以0的边界情况 TPR = TP / (TP + FN) if (TP + FN) != 0 else 0.0 FPR = FP / (FP + TN) if (FP + TN) != 0 else 0.0 return TPR, FPR
步骤3:遍历阈值生成ROC点
我们需要覆盖从0到1的所有可能阈值,比如按0.01的步长取点,这样能得到平滑的曲线:
import matplotlib.pyplot as plt # 把真实标签转换成0/1:裤子=1,套头衫=0 y_true = np.where(test_labels == "Trousers", 1, 0) # 取裤子的后验概率作为模型输出的正类分数 y_score = posteriors_test[:, 0] # 生成所有阈值 thresholds = np.arange(0.0, 1.01, 0.01) tpr_list = [] fpr_list = [] for thresh in thresholds: tpr, fpr = calculate_tpr_fpr(y_true, y_score, thresh) tpr_list.append(tpr) fpr_list.append(fpr) # 计算AUC(曲线下面积,可选但常用) def calculate_auc(fpr, tpr): # 按FPR排序确保递增 sorted_indices = np.argsort(fpr) fpr_sorted = np.array(fpr)[sorted_indices] tpr_sorted = np.array(tpr)[sorted_indices] # 梯形法计算面积 return np.trapz(tpr_sorted, fpr_sorted) auc_score = calculate_auc(fpr_list, tpr_list) # 绘制ROC曲线 plt.figure(figsize=(8,6)) plt.plot(fpr_list, tpr_list, label=f'ROC Curve (AUC = {auc_score:.2f})') plt.plot([0,1], [0,1], 'k--', label='Random Guess') plt.xlabel('False Positive Rate (FPR)') plt.ylabel('True Positive Rate (TPR)') plt.title('ROC Curve: Trousers vs Pullovers') plt.legend() plt.show()
二、多分类场景(10类)
你的思路完全正确:用**One-vs-Rest(一对多)**的策略,把每个类单独作为正类,其余所有类作为负类,分别绘制10条ROC曲线。
步骤1:准备数据
假设:
test_labels:测试集真实标签(0-9的数字,对应FMNIST的10类)posteriors_test:测试集每个样本的后验概率数组,形状是(n_samples, 10),每列对应一个类的后验概率
步骤2:遍历每个类绘制ROC曲线
plt.figure(figsize=(10,8)) # FMNIST的10类名称 class_names = ["T-shirt/top", "Trouser", "Pullover", "Dress", "Coat", "Sandal", "Shirt", "Sneaker", "Bag", "Ankle boot"] for class_idx in range(10): # 当前类作为正类,其余所有类作为负类 y_true = np.where(test_labels == class_idx, 1, 0) # 当前类的后验概率作为正类分数 y_score = posteriors_test[:, class_idx] # 计算当前类的TPR和FPR thresholds = np.arange(0.0, 1.01, 0.01) tpr_list = [] fpr_list = [] for thresh in thresholds: tpr, fpr = calculate_tpr_fpr(y_true, y_score, thresh) tpr_list.append(tpr) fpr_list.append(fpr) # 计算当前类的AUC auc_score = calculate_auc(fpr_list, tpr_list) # 绘制曲线 plt.plot(fpr_list, tpr_list, label=f'{class_names[class_idx]} (AUC={auc_score:.2f})') # 绘制随机猜测的基准线 plt.plot([0,1], [0,1], 'k--', label='Random Guess') plt.xlabel('False Positive Rate (FPR)') plt.ylabel('True Positive Rate (TPR)') plt.title('ROC Curves for 10-class FMNIST') plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left') plt.tight_layout() plt.show()
关键细节解释
- 阈值的作用:阈值是用来判断“当样本属于正类的概率大于等于这个值时,就预测为正类”。遍历所有可能的阈值,就能看到模型在不同严格程度下的表现——比如阈值设为0,所有样本都被预测为正类,此时TPR=1,FPR=1;阈值设为1,所有样本都被预测为负类,此时TPR=0,FPR=0。
- 为什么不用最大后验概率直接预测:ROC需要的是概率值而不是硬分类结果,这样才能展示模型在不同阈值下的权衡。你之前的“取最大后验概率”其实相当于一个隐含的阈值(比其他所有类的概率都高),但ROC要覆盖所有可能的阈值情况。
- 二值化特征不影响ROC:你的特征已经二值化了,但朴素贝叶斯计算出的后验概率是连续的(0到1之间),这正好适合做ROC分析——我们需要的就是这个连续的概率值来调整阈值。
内容的提问来源于stack exchange,提问作者POOJA GUPTA
相关产品推荐
相关产品推荐

