Skip to content

模型构建

本章目标

  1. 明确 train_model(...) 如何构建并训练 GaussianNB
  2. 理解 GaussianNB 的构造器参数 priorsvar_smoothing 的数学含义。
  3. 看清训练完成后最重要的模型属性及其与数学公式的对应关系。

重点方法与概念速览

名称类型作用
train_model(...)函数构建并训练一个 GaussianNB 模型,打印训练日志
GaussianNB(...)scikit-learn 提供的高斯朴素贝叶斯分类器——对连续特征的每个类别每个特征拟合 N(μkj,σkj2)
model.fit(X_train, y_train)方法在训练数据上统计类别先验和特征高斯参数——纯统计计算,无迭代优化
model.classes_属性模型识别到的类别标签数组
model.class_prior_属性各类别先验概率 P(Y=ck)
model.theta_属性各类别各特征的均值 μkj
model.var_属性各类别各特征的方差 σkj2(平滑后)

1. train_model(...) 的函数签名

参数速览

适用函数:train_model(X_train, y_train, var_smoothing=1e-9)

参数名类型说明示例取值
X_trainarray_like标准化后的训练特征矩阵,形状 (120,4),传入 GaussianNB.fit()X_train_s
y_trainarray_like训练标签向量,形状 (120,),取值 yi{0,1,2}y_train
var_smoothingfloat方差平滑项 ϵ。实际计算 σkj2+ϵσmax2,防止 σkj20 数值崩溃。默认 1e-91e-91e-8
返回值GaussianNB已完成 fit() 的模型对象,含 classes_class_prior_theta_var_ 等属性

示例代码

python
from model_training.classification.naive_bayes import train_model

model = train_model(X_train_s, y_train)

理解重点

  • 当前入口很直接:只负责构建一个 GaussianNBfit,没有变体对比或超参数搜索。
  • 所有默认超参数都写在函数签名里,阅读成本低,适合作为源码入门。
  • train_model(...) 是对 sklearn.naive_bayes.GaussianNB 的薄封装——算法本体在 sklearn,本仓库负责组织日志和工程流程。

2. GaussianNB 构造器参数

参数速览

适用 API:GaussianNB(priors=None, var_smoothing=1e-9)

参数名类型说明示例取值
priorsarray_likeNone类别的先验概率 P(Y=ck)None 时从训练数据估计:P(Y=ck)=nk/N。可手动传入数组覆盖数据估计None[0.3, 0.3, 0.4]
var_smoothingfloat方差平滑项 ϵ。最终方差为 σkj2+ϵσmax2,其中 σmax2 是所有特征所有类别中最大的方差。防止 σkj2012πσ2 导数炸1e-91e-81e-7

示例代码

python
from sklearn.naive_bayes import GaussianNB

model = GaussianNB(var_smoothing=1e-9)
model.fit(X_train_s, y_train)

理解重点

  • GaussianNB 的参数极简——一共两个:priors(先验)和 var_smoothing(数值保护)。这反映了朴素贝叶斯"参数少、假设强"的特点。
  • priors 默认为 None 时从训练数据按频率估计,对 iris 均衡数据来说三个类别的先验约为 [0.33,0.33,0.33]
  • var_smoothing 是当前分册最重要的超参数——它直接关联到方差为零时的数值稳定性问题。
  • GaussianNB 的 fit() 不涉及迭代优化——它只是扫描数据统计均值和方差。这与逻辑回归的 lbfgs 迭代和决策树的递归分裂形成鲜明对比。

3. 训练完成后的关键属性

参数速览

属性名类型数学含义说明
classes_ndarray,形状 (n_classes,){c1,c2,c3}模型识别到的类别标签列表,iris 中为 [0, 1, 2]
class_prior_ndarray,形状 (n_classes,)P(Y=ck)=nk/N各类别的先验概率
class_count_ndarray,形状 (n_classes,)nk训练集中各类别的样本数
theta_ndarray,形状 (n_classes, n_features)μkj各类别各特征的均值,对应高斯分布的位置参数
var_ndarray,形状 (n_classes, n_features)σkj2(平滑后)各类别各特征的方差,对应高斯分布的尺度参数——已应用 var_smoothing
epsilon_floatϵσmax2var_smoothing 对应的实际平滑绝对值

示例代码

python
print(f"类别: {model.classes_.tolist()}")
print(f"类别先验: {model.class_prior_.round(4)}")
print(f"各类别样本数: {model.class_count_}")
print(f"均值(theta_):\n{model.theta_}")
print(f"方差(var_):\n{model.var_}")

理解重点

  • theta_var_ 是 GaussianNB 最核心的两个训练产出——它们就是各类别下各特征高斯分布的参数。
  • theta_ 形状 (3,4) 意味着 3 个类别 × 4 个特征 = 12 个均值;var_ 同样有 12 个方差——模型一共只估计 24 个数字,训练极快。
  • class_prior_ 把"先验概率"这一理论概念直接映射为可观察的数值,是理解生成式分类思路的入口。
  • epsilon_ 提供了方差平滑的实际量级,对于理解 var_smoothing 是否真正生效有参考价值。

4. 训练阶段的工程封装

除了 GaussianNB(...).fit(...) 之外,train_model(...) 还做了几层工程包装:

参数速览

输出项作用
@print_func_info 标题在终端中定位训练入口
@timeit 训练耗时观察 fit() 的执行时间——对 GaussianNB 通常是毫秒级
var_smoothing 日志确认当前平滑参数配置
类别 日志确认多分类类别集合
类别先验 日志观察各类别基础比例,对应 P(Y=ck)

理解重点

  • 当前封装强调的是教学型可读性——通过装饰器打印函数信息和耗时,通过 print 输出关键属性。
  • 这一层把"构建模型""训练模型""打印结果"收在一个函数里,方便流水线和文档复用。
  • 从工程角度看,这样的拆分让 pipelines/classification/naive_bayes.py 保持简洁——编排层不需要关心日志打印细节。

常见坑

  1. 误以为当前实现使用的是所有朴素贝叶斯的通用封装——train_model 明确构建 GaussianNB,不是 MultinomialNBBernoulliNB
  2. 只知道 predict(...),却忽略 theta_var_class_prior_ 才是理解概率分类本质的关键属性。
  3. 忘记当前 X_train 应该是标准化后的特征——虽然 GaussianNB 不像逻辑回归那样对尺度敏感,但标准化影响方差估计的稳定性和 PCA 可视化。
  4. 把训练函数和后续评估逻辑混在一起理解——train_model 只负责训练主模型,不负责混淆矩阵、ROC 等诊断。

小结

  • train_model(...) 是本仓库 Naive Bayes 的核心训练入口,是对 GaussianNB 的薄封装。
  • GaussianNB 只有两个构造器参数(priorsvar_smoothing),属于参数最少的分类模型之一。
  • 训练完成后的关键属性:theta_(均值 μkj)、var_(方差 σkj2)、class_prior_(先验概率)——全部是解析计算,无迭代优化。