简介:一份面向数据挖掘与机器学习初学者的Java源码资源,聚焦树型朴素贝叶斯算法的实现与应用。该算法在经典朴素贝叶斯基础上引入决策树结构,通过信息增益等准则选择最优属性划分类别,能更灵活地处理多类问题。源码设计清晰,覆盖数据预处理、条件概率计算、决策树构建和分类预测等完整流程,适合用于文本分类、情感分析等场景。压缩包整体仅6KB,包含5个文件,其中4个Java源文件分别承担工具函数、属性互信息计算、树节点定义与主控逻辑,1个txt文件作为输入数据样例,便于直接运行验证。已有214人学习该资源,说明其实用性受到一定认可。通过研读源码,读者不仅能快速上手朴素贝叶斯变体的Java实现,还可掌握决策树与概率模型结合的关键细节,为后续深入研究人工智能算法打下基础。
1. 树型朴素贝叶斯:用一棵依赖树替换独立假设,Java实现的数据挖掘分类方案
朴素的“属性独立”伪命题,在树型朴素贝叶斯这里被松绑了一半:不再假设属性两两独立,而是用一棵依赖树表达属性之间的关联。这个权衡很有意思——它只比朴素贝叶斯多了一棵树,却能在很多数据挖掘分类任务里把精度拉高几个百分点;结构上又比完整贝叶斯网络简单得多,训练代价几乎可以忽略。这篇笔记讲的是这个算法的 Java 源码实现脉络:条件互信息怎么算、最大生成树怎么建、条件概率表怎么存、预测怎么做,以及我实际跑数据时才发现的几个坑。适合想在 Java 项目里落地一个可解释分类器、又在应付数据挖掘课程设计或面试题时被贝叶斯变种问住的开发者。
2. 从朴素贝叶斯到TAN:属性独立这一个假设,卡住了多少分类精度
2.1 朴素贝叶斯的三条软肋:公式之下藏着什么
朴素贝叶斯的分类决策写出来就一行:
P(c | x) ∝ P(c) · ∏ P(xi | c)
也就是说,给定类别 c,各属性 xi 之间完全独立。这个假设带来两个实际好处:参数估计只需要每个属性在各类别下的单变量分布,样本量要求低;训练是单趟扫描,内存和耗时都好控制。这也是为什么它至今仍是最常用的 baseline 模型,面试题里也总被拿来和 LR、树模型做对比。
代价也很直接。第一条软肋:当属性确实相关时,P(xi | c) 的连乘会系统性偏离真实联合概率。拿天气与运动场景举例,outlook=阴 和 humidity=高 在“去打球”这个类别下并不独立,阴天往往对应湿度偏高。连乘会把“阴天且湿度高”的概率算得过分低,一条本该判为不打球的样本被推给“去打球”,或者反过来,取决于偏置方向。这种偏差不是随机噪声,而是结构性失真,样本量再大也补不回来。
第二条软肋是冗余属性被重复加权。假设两个特征完全线性相关,朴素贝叶斯相当于把同一信息在 P(x1|c) 和 P(x2|c) 里各乘了一次,等于对这条证据给了双倍权重。特征越多,这种隐性加权越失控,甚至出现特征维度高到某个阈值后精度不升反降的现象。实际项目里,用户画像里有几十个强相关标签时,朴素贝叶斯的精度往往被同门的 LR 按在地上打。
第三条软肋更工程化:它要求属性天然离散,或者人工离散化。连续特征如果直接做密度估计塞进条件概率表,稀疏样本下的方差会大到离谱;而树型朴素贝叶斯的树结构学习同样需要离散特征支撑,这一点在 Java 实现里几乎躲不开。很多人第一次跑 TAN 源码翻车,就是栽在连续特征没做离散化上,这一条我在第 5 章还会展开。
这三条软肋不是“贝叶斯分类器不行”,而是“朴素”两个字带来的结构性限制。解决思路也顺理成章:把独立的图结构放宽成带依赖的图结构,但放宽的代价又不能太贵。于是就有了 TAN。
2.2 TAN 的树结构:每个属性最多多一个“帮手”
树型朴素贝叶斯(Tree-Augmented Naive Bayes,TAN)对图结构做了精确定义:所有属性节点都以类别节点为父节点;此外,每个属性节点最多再依赖一个其他属性节点。也就是说,把类别节点拿掉之后,属性之间的依赖子图是一棵有向树。
这种结构的数学表达是:
P(c | x) ∝ P(c) · P(x_root | c) · ∏ P(xi | parent(xi), c)
其中 parent(xi) 是属性 xi 在树上的唯一父属性,root 是树根属性,它只依赖类别。对比朴素贝叶斯,非根属性的条件概率从 P(xi|c) 换成了 P(xi|parent(xi), c),多了一个条件变量。预测时仍然走 argmax,不需要做任何图推理。
为什么刚好是一棵树?把结构放宽成树,新增的条件依赖数量级是 O(m),m 为属性数。每个非根属性只多一张二维概率表,参数总量从 O(m·k) 涨到 O(m·k·v),v 是父属性的取值数,通常是个位数到几十,完全可控。如果放宽成任意有向无环图,也就是完整贝叶斯网络,依赖边数量是 O(m²),结构学习要从搜索空间里暴力找,打分函数、禁忌搜索、模拟退火全都要上,训练开销直接上涨几个量级,而且小样本下极容易过拟合。
换句话说,TAN 是在“分类精度提升”和“训练代价基本不变”之间最划算的一档。相比之下,AODE 的思路是让所有属性两两配对,预测时对每个属性对求平均,存储成本和训练时间都高于 TAN,却并不能保证结构上更可解释。TAN 的 parent 关系是可以直接打印成一条依赖链给业务方看的,AODE 做不到这一点。
这里还有一个容易忽略的边界:如果属性之间的真实依赖关系是多层嵌套,比如 A 依赖 B、B 依赖 C、C 依赖 D,TAN 的“每个节点最多一个父属性”就装不下了。它只能近似成一条链,把最强的依赖关系挑出来。这个近似在大部分分类任务里够用,但如果你事先知道属性关系是深层漏斗状,TAN 不是最优选,直接上贝叶斯网络或者换成树模型更合适。
2.3 选型对比:TAN、朴素贝叶斯、AODE、贝叶斯网络怎么选
| 模型 | 依赖结构 | 参数规模 | 训练代价 | 适合场景 |
|---|---|---|---|---|
| 朴素贝叶斯 | 无属性依赖 | O(m·k) | 单趟扫描 | 属性基本独立、样本量小、要极快 baseline |
| TAN | 属性间一棵树 | O(m·k·v) | 互信息矩阵 O(m²) + 最大生成树 O(m²) | 属性有局部依赖、样本几千到几万、要可解释 |
| AODE | 所有属性对 | O(m²·k·v) | O(m²) 次统计扫描 | 样本量大、想用配对关系提升精度 |
| 完整贝叶斯网络 | 任意 DAG | 结构决定 | 结构搜索 NP 难 + 打分 | 属性关系由专家定义、不追求自动训练 |
实际项目里,我一般把 TAN 放在朴素贝叶斯之后作为第二个候选。先跑一个 NB 算出 baseline,再跑 TAN 对比提升:如果提升不足 1 个百分点,说明属性相关性影响确实有限,换 AODE 也未必有起色;如果提升超过 3 个百分点,说明独立假设已经明显失真,接下来值得试试带更多依赖的模型,甚至考虑换梯度提升树这类非参数模型。
如果你更熟悉 Python 那套数据挖掘生态,转过来看 Java 版实现,最大的差异在数据结构选择和循环写法上。Python 里 pandas 的一行 groupby 能做的统计,在 Java 里要手写 HashMap 计数,这也是源码读起来最费劲的地方。但反过来,Java 实现的好处是没有任何第三方依赖,一个工程文件就能跑,移植到 Hadoop 或 Spark 的 map 阶段也顺理成章。
3. 树的构建算法:条件互信息定权重,Prim算法找最大生成树
TAN 训练和朴素贝叶斯唯一的本质差别,是在参数估计之前要先把属性之间的依赖树“学”出来。这一步拆成三个子问题:用条件互信息度量属性关联强度;在完全图上跑最大生成树;确定根节点和边的方向。
3.1 条件互信息:度量“去掉类别影响后,两个属性还连不连”
两个属性之间的依赖强度,标准度量是条件互信息:
I(Xi; Xj | C) = Σ P(xi, xj, c) · log [ P(xi, xj | c) / (P(xi | c) · P(xj | c)) ]
直观理解:P(xi, xj | c) 是真实联合概率,P(xi | c) · P(xj | c) 是假设独立时的概率。两者相除取对数,衡量“在类别已知的前提下,知道 xi 之后对 xj 的预测增益有多大”,再对全空间加权求和。值越大,说明两个属性在类别把公共信息抽走之后仍然强相关,值得在树里连一条边;值接近 0,说明二者在类别已知后已无额外关联。
举一个手动可查的简化例子。假设只有三个属性 A、B、D,类别 C,某训练集算出的条件互信息矩阵是:
| 属性对 | I(Xi; Xj | C) |
|---|---|
| A-B | 0.12 |
| A-D | 0.42 |
| B-D | 0.31 |
那么 D 和 A 的关联最强,B 和 D 次之,A 和 B 几乎没有额外关联。最大生成树会先选 A-D 边(0.42),再选 B-D 边(0.31),得到 A-D-B 一条链;A-B 之间那 0.12 不会被选入,因为再加进去就会成环。
动手实现时,条件互信息的计算就是三层计数:联合计数 N(xi, xj, c)、成对计数 N(xi, c) 和 N(xj, c)、类别计数 N(c)。公式里的 log 项可以用任意底数,因为最大生成树只比较大小,不比较绝对值;但如果你习惯用 bit 报告数值,就除以 Math.log(2.0) 换底。需要注意的一点:条件互信息永远非负,不等于“越大越有线可挖”。小样本上它是一个有偏估计且偏高,这个问题我在 5.3 节给处理方案。
3.2 从完全图到最大生成树:为什么选Prim而不选Kruskal
有了互信息矩阵,下一步是在所有属性之间构建一棵生成树,使树上边的总权重最大。每个属性是图上的一个节点,任意两属性之间有一条边,权值就是互信息值——这是一个完全图,边数 E = m(m-1)/2。
最大生成树的两条经典路线是 Kruskal 和 Prim。Kruskal 把所有边按权重降序排序,逐个加入并查集,直到选出 m-1 条边;在边数多的时候,排序本身就是 O(E log E)。Prim 从任意节点出发,每轮选“连到当前树的最大边”的节点加入树,用邻接矩阵实现是严格的 O(m²)。
m 是属性个数,数据挖掘场景里通常是几十到几百这个量级。O(m²) 的 Prim 不需要排序,实现更短,而且在稠密完全图上是渐进最优的;Kruskal 虽然在稀疏图上表现好,但完全图没有稀疏性可言,构造边表再排序纯属多绕一圈。所以源码里直接用邻接矩阵版 Prim,工程上最省事。
还有一个工程细节:互信息矩阵是对称的,只需要存上三角,能省一半内存。属性数 200 时,double[][] 全量是 320KB,上三角是 160KB,差别不大;属性数 1000 时,全量要 8MB,上三角只要 4MB,这时候就有意义了。不过 Java 里二维数组的开销主要在对齐和对象头,我一般图省心直接全量存,m 超过 500 再见机行事。
算法走查一遍:初始化 0 号属性在树内,bestWeight[0] 给正无穷保证第一轮选中它;维护两个数组,bestWeight[v] 表示节点 v 能连到当前树上的最大边权,parent[v] 记录这条最大边连向树里的哪个节点;每一轮把 bestWeight 最大的未入树节点拉进树,并用它的边集去更新其他节点的 bestWeight 和 parent。跑完 m 轮,parent 数组就是生成树。
3.3 根节点与边的方向:训练时定下来,预测时才能查表
最大生成树是无向树,而 TAN 的推理需要方向:每个属性节点有唯一的父属性。标准做法是把类别节点当作根,从类别节点出发沿生成树做一次 BFS,给每条无向边定向:离开根的方向就是边的方向。距离类别节点最近的属性成为树的根属性,它只有类别节点这一个父节点;其余属性按层逐级挂靠。
把类别节点作为根不是随意选择。这样保证每个属性节点到类别节点的路径都尽量短,依赖方向与“类别驱动属性取值”这一生成直觉一致;同时也让预测时的每个条件概率都能在训练阶段直接查表——根属性查 P(xi|c),非根属性查 P(xi|parent(xi), c),预测时不需要做任何图上的概率推理或变量消元。
实现上,可以不显式建一棵有向树对象,只用两个数组就够。训练阶段拿到 Prim 输出的 parent 数组后,做一次从根属性的层序重定向,把无向父子关系转换成最终有向的 parentOfAttr 数组,根属性记 -1。预测阶段只需要查这个数组,每个属性走一步,总开销 O(m) 一次查表。
// 无向生成树 parent 数组 -> 有向依赖关系 parentOfAttr // 原始 parent[i] 只表示 i 入树时连到哪个节点,不知道谁是根 int m = treeParent.length; int[] parentOfAttr = new int[m]; int rootAttr = 0; // 这里取生成树起点为根属性,实际应从类别节点出发定层 Arrays.fill(parentOfAttr, -2); // -2 未处理 parentOfAttr[rootAttr] = -1; // -1 表示根属性,只依赖类别 // 从根属性开始做层序扩散,逐层确定方向 Queue<Integer> queue = new LinkedList<>(); queue.offer(rootAttr); while (!queue.isEmpty()) { int u = queue.poll(); for (int v = 0; v < m; v++) { if (treeParent[v] == u && parentOfAttr[v] == -2) { parentOfAttr[v] = u; // v 的父属性是 u queue.offer(v); } else if (treeParent[u] == v && parentOfAttr[v] == -2) { parentOfAttr[u] = v; // u 的父属性是 v queue.offer(v); } } }这段代码的核心是按 BFS 的层级把无向边统一改成“背离根”的方向。为什么不用 DFS?因为 BFS 天然保证先处理靠近根层的节点,便于逐层赋值,避免 DFS 在链式结构上递归过深导致栈溢出。属性数几百时 DFS 也没问题,但 BFS 更稳,而且代码可读性好。
4. Java源码实现:条件互信息、最大生成树与分类器的完整代码脉络
4.1 源码包结构:六个类管住整个TAN训练与预测流程
一个可维护的 TAN 实现,不需要把逻辑都塞进一个类。按数据加载、离散化、核心算法、概率存储四层拆,我一般这样组织工程:
tan-bayes/ ├── pom.xml // Maven工程,JDK 8+ └── src/main/java/com/mining/tan/ ├── core/ │ ├── TanBayesClassifier.java // 训练入口:互信息矩阵→生成树→概率表 │ ├── ConditionalMutualInfo.java// 条件互信息估计,只依赖数据矩阵 │ ├── SpanningTreeBuilder.java // Prim最大生成树,纯静态方法 │ └── ProbabilityTable.java // 条件概率表存储与拉普拉斯平滑查询 ├── data/ │ ├── DataLoader.java // ARFF/CSV加载成int[][]离散矩阵 │ └── Discretizer.java // 连续特征等频分箱,返回箱边界 └── util/ └── MathUtil.java // 对数累加、argmax等公共函数TanBayesClassifier 是唯一对外暴露训练/预测接口的门面类。训练流程是:DataLoader 读入原始数据,Discretizer 把连续列切成离散整数编码,然后依次调用 ConditionalMutualInfo 填满互信息矩阵、SpanningTreeBuilder 得到生成树、重定向成最终依赖关系、最后用 ProbabilityTable 逐属性建条件概率表。训练完成后,状态只保留三个对象:parentOfAttr 数组、classPrior 数组、tables 数组。中间的所有计数 Map 在训练结束后都可被 GC 回收。
这个划分的边界值得说一句:ConditionalMutualInfo 和 SpanningTreeBuilder 都是纯函数式类,不持有状态,方便单元测试;ProbabilityTable 是唯一存储模型参数的类,未来如果要接 PMML 导出,只需要增加一个序列化方法,不牵动别的类。把“统计计数”和“模型存储”分开,是我读很多数据挖掘源码之后养成的习惯。
4.2 条件互信息计算:Java实现与参数说明
核心方法只处理离散化的整数矩阵,每一行是样本,每一列是属性。这里用了一个小技巧:用位运算打包联合计数的键,完全避开字符串拼接,在几十万样本时节省非常明显。
public class ConditionalMutualInfo { /** * 计算属性 a1 与 a2 在类别 classIdx 条件下的互信息。 * * @param data 离散化后的训练数据,每行一个样本 * @param a1 第一个属性的列索引 * @param a2 第二个属性的列索引 * @param classIdx 类别列的索引 * @return I(a1; a2 | class),自然对数底,非负 */ public double compute(int[][] data, int a1, int a2, int classIdx) { int total = data.length; // 键用位打包:低16位存类别,中间16位存a2,高32位存a1 Map<Long, Integer> jointCount = new HashMap<>(); Map<Integer, Integer> aiGivenClassCount = new HashMap<>(); Map<Integer, Integer> ajGivenClassCount = new HashMap<>(); Map<Integer, Integer> classCount = new HashMap<>(); for (int[] row : data) { int c = row[classIdx]; int xi = row[a1]; int xj = row[a2]; jointCount.merge((((long) xi) << 32) | (((long) xj) << 16) | c, 1, Integer::sum); aiGivenClassCount.merge((c << 16) | xi, 1, Integer::sum); ajGivenClassCount.merge((c << 16) | xj, 1, Integer::sum); classCount.merge(c, 1, Integer::sum); } double mi = 0.0; for (Map.Entry<Long, Integer> e : jointCount.entrySet()) { long key = e.getKey(); int xi = (int) (key >>> 32); int xj = (int) ((key >> 16) & 0xFFFF); int c = (int) (key & 0xFFFF); int nJoint = e.getValue(); int nXiC = aiGivenClassCount.getOrDefault((c << 16) | xi, 0); int nXjC = ajGivenClassCount.getOrDefault((c << 16) | xj, 0); int nC = classCount.get(c); double pJoint = (double) nJoint / nC; double pXiC = (double) nXiC / nC; double pXjC = (double) nXjC / nC; mi += ((double) nJoint / total) * Math.log(pJoint / (pXiC * pXjC)); } return mi; } }参数说明:data 必须是整数编码的离散矩阵,类别列和其他属性列不能混用索引;a1、a2 不能等于 classIdx,调用方要提前过滤对角线。位运算的 16 位分割要求属性值和类别值都小于 65535,常规数据集没问题,属性取值上十万的文本特征必须先做分箱压缩。返回的是自然对数底的互信息,构建树只看相对大小;需要 bit 单位时,把返回值除以 Math.log(2.0)。时间复杂度是遍历一次数据矩阵,O(total × m),通常只占训练总耗时里很小一部分。
4.3 Prim最大生成树:把互信息矩阵变成parent数组
互信息矩阵是对称的,SpanningTreeBuilder 直接吃 double[][],输出一个 parent 数组。输出数组的含义是“无向生成树上的父子关系”,方向约定由调用方根据根节点重定向。
public class SpanningTreeBuilder { /** * 用 Prim 算法构建最大带权生成树。 * 复杂度 O(m^2),m 为属性个数。 * * @param matrix 对称条件互信息矩阵,matrix[i][j] = I(i; j | class) * @return parent 数组,parent[i] 是 i 在无向树上连接到的节点 */ public static int[] primMaxSpanningTree(double[][] matrix) { int m = matrix.length; boolean[] inTree = new boolean[m]; double[] bestWeight = new double[m]; int[] parent = new int[m]; Arrays.fill(bestWeight, Double.NEGATIVE_INFINITY); // 从 0 号属性开始生长 bestWeight[0] = Double.POSITIVE_INFINITY; parent[0] = 0; for (int round = 0; round < m; round++) { int u = -1; for (int v = 0; v < m; v++) { if (!inTree[v] && (u == -1 || bestWeight[v] > bestWeight[u])) { u = v; } } inTree[u] = true; for (int v = 0; v < m; v++) { if (!inTree[v] && matrix[u][v] > bestWeight[v]) { bestWeight[v] = matrix[u][v]; parent[v] = u; } } } return parent; } }逻辑说明:每一轮选出的 u 是“当前不在树内,但到树的连接边权最大”的节点,因此它入树时连接的节点一定是最优的。外层 m 轮、内层两次各 m 长度的扫描,总复杂度 O(m²)。这个实现假设 matrix 对称,如果上游互信息计算有数值误差导致不对称,最大生成树仍能跑,只是结果可能受轻微影响。排查时可以用 matrix[u][v] 与 matrix[v][u] 的差做一个断言,差超过 1e-10 就报警,通常能抓到索引传反的 bug。
4.4 概率表存储与预测:log域下避免下溢
训练阶段最后一步是逐属性填充条件概率表。根属性存二维表 [class][value];非根属性存三维表 [class][value][parentValue],每个格子在计数后加拉普拉斯平滑。预测阶段统一走 log 累加,避免几十个小概率连乘直接下溢成 0.0。
public class TanBayesClassifier { private int[] parentOfAttr; // -1 表示根属性,否则存父属性索引 private double[] classPrior; // P(c),经平滑 private ProbabilityTable[] tables; /** * 对一条离散样本做预测。 * * @param instance 属性值数组,长度必须等于训练时的属性数 * @return 后验概率最大的类别编码 */ public int predict(int[] instance) { int k = classPrior.length; double[] logPosterior = new double[k]; for (int c = 0; c < k; c++) { double logP = Math.log(classPrior[c]); for (int i = 0; i < instance.length; i++) { if (parentOfAttr[i] == -1) { logP += Math.log(tables[i].getProb(c, instance[i])); } else { int p = parentOfAttr[i]; logP += Math.log(tables[i].getProb(c, instance[i], instance[p])); } } logPosterior[c] = logP; } // argmax,不需要归一化,因为 log 域里统一减去常数不影响相对大小 int best = 0; for (int c = 1; c < k; c++) { if (logPosterior[c] > logPosterior[best]) { best = c; } } return best; } }注意 predict 里没有做概率归一化,argmax 只需要比较相对大小,log 域里加同一个常数不影响结果。如果业务上需要输出置信度,可以对 logPosterior 做 softmax 还原成归一化概率。tables[i].getProb 系列方法内部要处理下标越界:instance[p] 的取值必须落在训练时见过的父属性取值空间内,Discretizer 在预处理时要用训练集的箱边界切分新样本,而不是重新分箱,否则预测时查表下标会错位。多线程预测时,TanBayesClassifier 本身是无状态只读的,可以安全并发调用 predict。
5. 树型朴素贝叶斯避坑指南:五个训练与预测阶段的翻车现场
TAN 的实现看起来不长,但真正跑数据的阶段,问题几乎都集中在下面五个地方。每一条都是“现象→原因→解决”的结构,按我遇到的出现频率排序。
5.1 零概率陷阱:没见过的组合让整条样本被否决
现象:训练集里某些条件组合没有出现,预测时对应概率是 0,连乘后整个类别的后验变成 0,样本莫名其妙被分到另一个类别。更隐蔽的是,如果两个类别的后验都是 0,argmax 会退化成随机返回第一个类别,线上表现完全不可控。
原因:条件概率表是按有限样本统计的,离散属性取值组合数一多,稀疏组合必然出现。TAN 比朴素贝叶斯更容易踩中,因为非根属性的表是二维条件,组合数量是根属性表的 v 倍。类别多、属性多、每个属性取值多的时候,零概率几乎是必然事件。
解决:拉普拉斯平滑是标准手段。根属性 P(xi|c) = (Nxi_c + α) / (Nc + α·k),非根属性 P(xi|parent,c) = (Nxi_parent_c + α) / (Nparent_c + α·k)。α 取 1 是最常见选择,等价于加 1 平滑;数据量小时可以试 0.5,数据量大时 0.1 甚至更小。千万不要把 α 设到 10 或以上,所有概率会被拉成均匀分布,树结构带来的精度提升直接被抹平。这个 α 建议作为构造参数暴露出来,方便调参。
5.2 连续属性裸奔:不离散化就训练,树结构全是假连接
现象:直接把体温、金额这类连续值塞进互信息计算,得到的结果对阈值极其敏感。换一条样本、阈值微变,互信息值和树结构就完全不同;预测更是毫无稳定性,同一模型跑两次推理结果不一致。
原因:TAN 的概率表是离散编码的,连续值的“取值集合”无穷大,统计计数形同虚设。条件互信息理论本身可以定义在连续变量上,但工程实现里不会真的去搞数值积分。很多源码为了提高运行速度,直接用 int 型二维数组存数据,连续值一旦被强转成 int,等于随机分箱,信息损失不可控。
解决:训练前强制离散化。等频分箱(每个箱样本量尽量均匀)通常好于等宽分箱,因为等宽遇到长尾分布时大部分箱里样本极少。箱数我一般取 5 到 10,取太少丢信息,取太多零概率爆炸。关键细节:离散化器只在训练集上拟合,预测时用训练集的箱边界切分新样本。很多人在这里翻车是因为每次预测都对全量数据重新分箱,训练和预测的编码空间不对齐,查表全错位。
5.3 小样本下条件互信息虚高:噪声连接进树,“帮忙”变捣乱
现象:样本量几百时构建出的树里经常出现两条明显不该相连的属性,比如“用户注册天数”和“支付金额”挂在一起。交叉验证发现,去掉这根边后精度反而更高。
原因:条件互信息是渐近无偏估计,样本有限时估计值带正偏差。属性取值组合越多,偏差越大,因为联合计数 nJoint 被散得很稀疏,log 项的分子分母都失去稳定性,个别样本就能把互信息顶得很大。真实信号被噪声淹没后,最大生成树选出的边就不一定代表真实依赖。
解决:常用做法是给互信息矩阵做显著性校正。对每对属性做随机置换检验,计算 p 值,只保留 p 值小于 0.05 的边,其余边权值直接视为 0。训练集每个类别低于 100 条样本时,这一步建议必开;样本量上万后,估计已经稳定,置换检验的收益就很小了。另一种折中是直接在互信息估计时给联合计数加一个小先验,效果等价于把 log 项的分母往上托一点,实现简单,但不如置换检验有统计依据。
5.4 概率连乘下溢:小概率项的连乘积在double里直接变零
现象:属性 20 个、每个条件概率都不超过 0.1 时,连乘后数值低于 Double.MIN_VALUE,log 后验出现 -Infinity。如果所有类别都变成 -Infinity,argmax 失效,预测结果等于随机。
原因:贝叶斯派模型的通病。条件概率连乘是乘积形式,数值范围指数级缩小。double 虽然能表示很小的非规格化数,但几十项连乘仍然会触底。这个问题在朴素贝叶斯里就有,TAN 因为多了一层条件概率,数值反而更小,触底更快。
解决:predict 里必须用 log 域,把连乘改成连加。如果某些场景必须输出真实的归一化后验概率,对 log 后验做一次 softmax,先减去最大值再做指数累加,得到归一化分母。另有一个细节:不要用 Math.pow 做连乘,那是数值灾难的加速器,中间过程毫无必要地放大了精度损失。
5.5 类别不平衡:先验概率把后验全拉偏
现象:二分类里正负比 9:1,模型精度看着有 90%,但看混淆矩阵,负类几乎全被吞成正类。TAN 在这种数据上比朴素贝叶斯不见得好,因为树的构建不受先验影响,但预测受先验影响很大。
原因:classPrior 里多数类先验大,乘到后验后把少数类的条件概率优势盖掉。树结构学到的依赖关系再准,也架不住先验这个偏置。尤其当少数类本身条件概率略有波动时,先验差一个数量级,后验直接翻盘。
解决:两档处理。轻量做法是训练后把 classPrior 改成均匀分布,或按业务赔率加权;正规做法是训练时对少数类过采样或对多数类降采样,再把处理后的先验估计回原始比例。实践里我建议先跑一次均匀先验的预测,看条件概率本身是否可分,再决定要不要动样本。如果均匀先验下少数类仍然被压制,问题不在先验而在特征,调结构比调先验有用。
6. 验证实现与进阶:对称性检查、交叉验证和结构打印
写完这一套代码,先别急着上生产。我验证 TAN 实现有一个固定套路,三步能拦住绝大多数隐性 bug。
第一步是验证互信息矩阵的数学性质:输出一轮全属性对的互信息值,断言每对 (i,j) 和 (j,i) 对称、非负、对角线为 0。这个检查能拦住上游数据错位、索引混用一类问题。第二步是 K 折交叉验证,和朴素贝叶斯对照跑同一份离散化数据,TAN 通常能高 1 到 5 个百分点;如果完全没有提升,优先怀疑树结构构建环节,把生成树打出来看看,边是不是连在互信息最大的两两之间。直接打印 parentOfAttr 数组是最快的排障方式:
// 打印属性依赖树,肉眼判断是否存在不合理的连边 for (int i = 0; i < parentOfAttr.length; i++) { int p = parentOfAttr[i]; if (p == -1) { System.out.println("attr " + i + " (root)"); } else { System.out.println("attr " + i + " -> parent attr " + p); } }第三步是留一校验做细粒度兜底。样本少的项目直接留一交叉验证,样本多就用 5 折。如果发现某几折精度波动特别大,回头看该折训练集里是否恰好把某个取值组合整体抽走了——这往往是零概率平滑参数需要调大的信号。
进阶方向上,最有性价比的是把离散化从等频换成 MDL 分箱,让箱边界跟着类别分布走,能再挤出一两个百分点;再往后就是把 TAN 作为基分类器塞进 bagging 框架里,用多棵树的随机性弥补单棵结构的刚性。这套 Java 实现不需要任何第三方依赖,核心就是互信息、Prim 和概率表三件事,把这三件事吃透,TAN 也就成了数据挖掘源码里一块很规整的积木。我个人的习惯是永远先跑对称性再看精度,结构对了结果自然能解释,结构错了精度再高也不敢上线。希望帮到你。
本文还有配套的精品资源,点击获取