Skip to content

朴素贝叶斯

朴素贝叶斯(Naïve Bayes)是一种基于贝叶斯定理的概率分类算法。之所以称其为“朴素”(naïve),是因为它做出了一个很强的(且常常不切实际的)假设:给定类别标签的情况下,所有预测特征(predictor features)之间相互独立。

尽管存在这种简化,朴素贝叶斯分类器在许多实际场景中表现出人意料的好,特别是在文本分类(如垃圾邮件过滤)和医疗诊断领域。

其核心思想是利用贝叶斯定理计算给定特征下,数据点属于某个特定类别的概率:

P(Class | Features) = [P(Features | Class) * P(Class)] / P(Features)

其中:

  • P(Class | Features): 后验概率 (Posterior Probability) - 在观测到特征后,数据点属于某个类别的概率(这是我们希望求得的)。
  • P(Features | Class): 似然 (Likelihood) - 在给定类别下,观测到这些特征的概率。正是在这一步,“朴素”独立性假设简化了计算:P(feature1, feature2, ... | Class) ≈ P(feature1 | Class) * P(feature2 | Class) * ...
  • P(Class): 先验概率 (Prior Probability) - 数据点属于某个类别的总体概率(基于该类别在训练数据中的频率)。
  • P(Features): 证据 (Evidence) - 观测到这些特征的总体概率。由于对于给定的数据点,这个值对于所有类别都是常数,在分类时常常被忽略(我们只需比较分子)。

该算法计算每个类别的后验概率,并将数据点分配给具有最高概率的类别。

Scikit-learn 提供了几种朴素贝叶斯实现,它们的主要区别在于对 P(Features | Class) 分布的假设:

  • 高斯朴素贝叶斯 (GaussianNB): 假设特征(给定类别下)服从高斯(正态)分布。适用于连续数值型特征。
  • 多项式朴素贝叶斯 (MultinomialNB): 假设特征代表计数或频率(例如,文本分类中的词频)。适用于离散特征。
  • 伯努利朴素贝叶斯 (BernoulliNB): 假设特征是二元的(0 或 1,表示出现或不出现)。常用于使用二元词项出现特征的文本分类。
  • 补集朴素贝叶斯 (ComplementNB): 是 MultinomialNB 的一种改进,特别适用于不平衡数据集。
  • 分类朴素贝叶斯 (CategoricalNB): 专为分类分布的特征设计。

选择哪种分类器取决于你的特征类型。

使用 Python 构建高斯朴素贝叶斯模型

Section titled “使用 Python 构建高斯朴素贝叶斯模型”

让我们使用 Scikit-learn 实现一个高斯朴素贝叶斯分类器。我们将生成一些合成数据进行演示。

import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.datasets import make_blobs
from sklearn.model_selection import train_test_split
from sklearn.naive_bayes import GaussianNB
from sklearn.metrics import accuracy_score, confusion_matrix, classification_report
# 1. 生成合成数据 (具有高斯分布的 blob)
# 使用 2 个特征便于可视化
X, y = make_blobs(n_samples=300, centers=2, n_features=2,
random_state=42, cluster_std=1.5)
# 2. 将数据分割为训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
# 3. 初始化并训练 GaussianNB 模型
model_gnb = GaussianNB()
model_gnb.fit(X_train, y_train)
# 4. 在测试集上进行预测
y_pred = model_gnb.predict(X_test)
# 5. 评估模型
accuracy = accuracy_score(y_test, y_pred)
conf_matrix = confusion_matrix(y_test, y_pred)
class_report = classification_report(y_test, y_pred)
print(f"Accuracy: {accuracy:.4f}\n")
print("Confusion Matrix:")
print(conf_matrix)
print("\nClassification Report:")
print(class_report)
# 6. 可选:可视化决策边界 (对于二维数据)
plt.figure(figsize=(10, 6))
# 绘制训练数据点
sns.scatterplot(x=X_train[:, 0], y=X_train[:, 1], hue=y_train, palette='viridis', marker='o', label='Train Data')
# 创建用于绘制决策边界的网格
ax = plt.gca()
xlim = ax.get_xlim()
ylim = ax.get_ylim()
xx, yy = np.meshgrid(np.linspace(xlim[0], xlim[1], 100),
np.linspace(ylim[0], ylim[1], 100))
Z = model_gnb.predict_proba(np.c_[xx.ravel(), yy.ravel()])[:, 1]
Z = Z.reshape(xx.shape)
# 绘制决策边界和间隔
plt.contourf(xx, yy, Z, levels=[0, 0.5, 1], cmap='viridis', alpha=0.3)
plt.contour(xx, yy, Z, levels=[0.5], colors='black', linestyles='--')
plt.title('Gaussian Naive Bayes Decision Boundary')
plt.xlabel('Feature 1')
plt.ylabel('Feature 2')
plt.legend(loc='upper right')
plt.grid(True)
plt.show()
# 7. 可选:显示部分测试点的后验概率
probabilities = model_gnb.predict_proba(X_test)
print("\nPosterior Probabilities for first 5 test points (Class 0, Class 1):")
print(np.round(probabilities[:5], 3))
Accuracy: 1.0000
Confusion Matrix:
[[45 0]
[ 0 45]]
Classification Report:
precision recall f1-score support
0 1.00 1.00 1.00 45
1 1.00 1.00 1.00 45
accuracy 1.00 90
macro avg 1.00 1.00 1.00 90
weighted avg 1.00 1.00 1.00 90
Posterior Probabilities for first 5 test points (Class 0, Class 1):
[[1. 0. ]
[0. 1. ]
[1. 0. ]
[0. 1. ]
[0. 1. ]]

输出显示,在该简单且分离良好的合成数据集上获得了完美准确率。混淆矩阵确认没有错误分类。分类报告提供了每个类别的精度(precision)、召回率(recall)和 F1 分数。图表可视化了数据点以及高斯朴素贝叶斯分类器找到的线性决策边界。后验概率显示了模型对前几个测试样本预测的置信度。

  • 快速高效: 训练和预测速度非常快,即使在大型数据集上也是如此。
  • 所需数据较少: 即使训练数据量较小,也能表现得相当好。
  • 可扩展性强: 随特征数量和数据点的数量呈线性扩展。
  • 处理高维数据: 在特征数量较多时(如文本分类)表现良好。
  • 概率性: 自然地提供预测的概率估计。
  • 处理不同特征类型: 不同的变体可以处理连续、离散和二元特征。
  • 适合多类别问题: 本质上适用于多类别分类问题。
  • “朴素”独立性假设: 特征独立的这个核心假设在实际数据中经常不成立,这有时会限制模型的性能。
  • “零频率”问题: 如果某个特定的特征值或类别在训练数据中没有出现在某个给定类别下,模型可能会给它分配零概率,这可能会主导整个计算结果。可以使用拉普拉斯平滑(加法平滑)等技术来缓解这个问题,这些技术通常在库中默认实现。
  • 对特征分布敏感: 性能取决于假设的特征分布(高斯、多项式等)与实际数据分布的匹配程度。

尽管朴素贝叶斯算法简单,但它在各种应用中非常有效:

  • 文本分类: 垃圾邮件过滤(识别垃圾邮件)、情感分析(将文本分类为积极/消极/中立)、文档分类。
  • 医疗诊断: 根据症状(特征)预测疾病的概率。
  • 实时预测: 其速度使其适用于需要快速预测的应用。
  • 推荐系统: 可以作为混合推荐系统的一部分。