一、贝叶斯分类基础知识

1.1 核心思想:贝叶斯定理

贝叶斯分类的核心是基于贝叶斯定理,通过计算样本属于不同类别的后验概率,选择概率最大的类别作为预测结果。贝叶斯定理的公式如下:

各符号含义:

P(Y|X) :后验概率,已知样本特征X时,样本属于类别Y的概率(这是我们最终要计算的)

P(Y) :先验概率,样本属于类别Y的初始概率(可通过训练集中各类别样本占比估算)

P(X|Y):似然概率,在类别Y的条件下,样本具有特征X的概率

( P(X) \):证据因子,样本特征X出现的概率(对所有类别而言是固定值,计算时可忽略,不影响类别判断)

1.2 朴素贝叶斯的“朴素”之处

朴素贝叶斯分类器在贝叶斯定理的基础上,增加了一个关键假设:样本的所有特征之间相互独立这个假设大大简化了计算难度——原本计算 \( P(X|Y) \) 需要考虑所有特征的联合概率,而特征独立假设下,联合概率可分解为各个特征条件概率的乘积:

其中n是样本的特征数量, Xi 是样本的第i个特征。虽然“特征独立”的假设在现实场景中很难完全成立,但朴素贝叶斯分类器依然能在很多任务中表现出色,且具有计算高效、泛化能力强的优点。

1.3 常见的朴素贝叶斯模型

根据样本特征的分布类型,朴素贝叶斯分为不同的实现版本,实验中主要用到以下两种:

  • 高斯朴素贝叶斯(GaussianNB):适用于特征是连续值的场景,假设每个类别的特征都服从高斯分布(正态分布);

  • 多项式朴素贝叶斯(MultinomialNB):适用于特征是离散值(如计数、频率)的场景,假设特征服从多项式分布(如文本分类中的词频特征)。

二、实战贝叶斯模型

3.1 模块1:环境配置与数据加载

功能说明:导入实验所需库,加载鸢尾花(连续特征)和手写数字(离散特征)数据集,查看数据基本信息(验证数据集完整性、特征类型、类别分布)。

# 导入核心库
import numpy as np
from sklearn.datasets import load_iris, load_digits
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
from sklearn.naive_bayes import GaussianNB, MultinomialNB
from sklearn.metrics import accuracy_score, confusion_matrix, classification_report

# 加载鸢尾花数据集(连续特征)
iris = load_iris()
X_iris, y_iris = iris.data, iris.target  # X_iris: 特征矩阵(150,4),y_iris: 标签(150,)
print("鸢尾花数据集信息:")
print(f"特征数量:{X_iris.shape[1]},样本数量:{X_iris.shape[0]}")
print(f"类别:{np.unique(y_iris)}(对应品种:{iris.target_names})\n")

# 加载手写数字数据集(离散特征:像素灰度值0-16)
digits = load_digits()
X_digits, y_digits = digits.data, digits.target  # X_digits: 特征矩阵(1797,64),y_digits: 标签(1797,)
print("手写数字数据集信息:")
print(f"特征数量:{X_digits.shape[1]},样本数量:{X_digits.shape[0]}")
print(f"类别:{np.unique(y_digits)}(0-9手写数字)")

3.2 模块2:数据预处理

功能说明:对数据进行划分(训练集/测试集)和标准化(仅连续特征),避免量纲影响,保证实验可复现。

# 3.2.1 数据划分:按7:3比例拆分训练集和测试集,random_state保证可复现
# 鸢尾花数据集(连续特征)
X_iris_train, X_iris_test, y_iris_train, y_iris_test = train_test_split(
    X_iris, y_iris, test_size=0.3, random_state=42  # test_size=0.3表示30%为测试集
)

# 手写数字数据集(离散特征)
X_digits_train, X_digits_test, y_digits_train, y_digits_test = train_test_split(
    X_digits, y_digits, test_size=0.3, random_state=42
)

# 3.2.2 特征标准化:仅对连续特征(鸢尾花)进行,离散特征无需标准化
scaler = StandardScaler()
X_iris_train_scaled = scaler.fit_transform(X_iris_train)  # 训练集拟合+标准化
X_iris_test_scaled = scaler.transform(X_iris_test)  # 测试集仅标准化(复用训练集参数)

print("数据预处理完成:")
print(f"鸢尾花训练集样本数:{X_iris_train_scaled.shape[0]},测试集样本数:{X_iris_test_scaled.shape[0]}")
print(f"手写数字训练集样本数:{X_digits_train.shape[0]},测试集样本数:{X_digits_test.shape[0]}")

3.3 模块3:模型构建与训练

功能说明:根据数据集特征类型选择对应朴素贝叶斯模型,使用训练集训练模型。

# 3.3.1 高斯朴素贝叶斯(适配鸢尾花连续特征)
gnb = GaussianNB()
gnb.fit(X_iris_train_scaled, y_iris_train)  # 用标准化后的训练集训练

# 3.3.2 多项式朴素贝叶斯(适配手写数字离散特征)
mnb = MultinomialNB()
mnb.fit(X_digits_train, y_digits_train)  # 离散特征直接训练

print("两个模型训练完成!")

3.4 模块4:模型预测与性能评估

功能说明:使用训练好的模型对测试集进行预测,通过准确率、混淆矩阵、分类报告评估模型性能。

# 3.4.1 鸢尾花模型预测与评估
y_iris_pred = gnb.predict(X_iris_test_scaled)  # 测试集预测
print("="*50)
print("鸢尾花数据集分类结果评估:")
print(f"准确率:{accuracy_score(y_iris_test, y_iris_pred):.4f}")  # 计算准确率
print("\n混淆矩阵:")
print(confusion_matrix(y_iris_test, y_iris_pred))  # 混淆矩阵(行:真实标签,列:预测标签)
print("\n分类报告(精确率/召回率/F1分数):")
print(classification_report(y_iris_test, y_iris_pred, target_names=iris.target_names))  # 详细评估

# 3.4.2 手写数字模型预测与评估
y_digits_pred = mnb.predict(X_digits_test)  # 测试集预测
print("="*50)
print("手写数字数据集分类结果评估:")
print(f"准确率:{accuracy_score(y_digits_test, y_digits_pred):.4f}")
print("\n混淆矩阵:")
print(confusion_matrix(y_digits_test, y_digits_pred))
print("\n分类报告(精确率/召回率/F1分数):")
print(classification_report(y_digits_test, y_digits_pred))

3.5 模块5:参数影响验证

功能说明:验证多项式朴素贝叶斯alpha平滑系数对模型性能的影响。

# 测试不同alpha值(平滑系数)对多项式朴素贝叶斯的影响
alpha_values = [0.1, 1.0, 10.0]
print("="*50)
print("不同alpha值对於手写数字分类准确率的影响:")
for alpha in alpha_values:
    mnb_temp = MultinomialNB(alpha=alpha)
    mnb_temp.fit(X_digits_train, y_digits_train)
    y_temp_pred = mnb_temp.predict(X_digits_test)
    acc = accuracy_score(y_digits_test, y_temp_pred)
    print(f"alpha={alpha}时,准确率:{acc:.4f}")

三、实验结果与分析

3.1 基础实验结果

运行上述完整代码后,预期得到以下核心结果:

  • 鸢尾花数据集(高斯NB):准确率约97.78%,混淆矩阵显示仅1个Versicolor品种被误分为Virginica,三个品种的精确率、召回率、F1分数均≥0.96;

  • 手写数字数据集(多项式NB):准确率约88.33%,混淆矩阵显示数字3、8、9误分类较多,数字0、1、6的F1分数≥0.95,数字3、8的F1分数≤0.8。

3.2 结果分析

3.2.1 数据集特征与模型适配性

鸢尾花数据集特征区分度高、相关性低,较好满足“特征独立”假设,且连续特征适配高斯分布,因此高斯NB性能优异;手写数字数据集特征为像素灰度值(离散),适配多项式NB,但部分数字(3与8、9与4)像素分布相似,破坏了特征独立假设,导致误分类增多。

3.2.2 参数影响分析

alpha平滑系数的作用是避免似然概率为0:

  • alpha=0.1:准确率约89.26%,平滑不足,对少数样本拟合更好,相似数字分类效果略有提升;

  • alpha=1.0(默认):准确率88.33%,平滑适中,平衡拟合效果;

  • alpha=10.0:准确率约86.11%,过度平滑导致模型欠拟合,丢失关键特征信息,性能下降。

 

更多推荐