多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

朴素贝叶斯二分类器手写实现:平滑与对数概率技巧

朴素贝叶斯二分类器手写实现:平滑与对数概率技巧 美团春招算法岗突然考了一道朴素贝叶斯二分类器当时我盯着题目愣了几秒——这玩意不是机器学习课上的基础模型吗笔试居然也会手写等冷静下来才发现这道题恰恰是很多人的分水岭原理大家都懂真要在限时里写出一个干净、正确、能跑通的二分类器考验的其实是概率统计的落地能力、编码熟练度和对细节的把控。这篇文章我会把题目、思路、三种语言的完整实现和在线测试方案都拆开讲清楚。不论你是准备算法岗秋招/春招还是刚开始接触机器学习手动实现这篇都可以直接当作一份可复现的参考。1. 题目拆解与考点分析1.1 我复现的题目版本先说题面。由于是回忆版本我把输入输出格式重新整理过一轮确保三种语言实现起来一致。题目朴素贝叶斯二分类器给定n个训练样本每个样本由m个离散特征和 1 个二分类标签组成。特征取值均为非负整数标签为 0 或 1。要求训练一个朴素贝叶斯分类器对q个测试样本分别输出预测标签。输入描述第一行三个整数n,m,q含义同上。接下来n行每行m1个整数前m个是特征值最后一个是标签。接下来q行每行m个整数表示一个待预测样本。输出描述输出q行每行一个整数 0 或 1。数据约定1 ≤ n ≤ 10001 ≤ m ≤ 201 ≤ q ≤ 100。特征取值范围0 ≤ x ≤ 10^9。测试集中任意样本的某个特征列可能从未在训练集中出现过。这个约定最后一条是核心它意味着你不能只保存训练集出现过的取值否则预测时会遇到“条件概率为 0”必须用拉普拉斯平滑兜底。这个细节和线上笔试的隐藏样例高度相关后面会单独展开。1.2 这道题到底在考什么从面试官角度手写朴素贝叶斯二分类器至少有四个明确的考察点第一概率基础是否扎实。贝叶斯公式、条件独立性假设、先验概率和后验概率的关系这些概念如果只停留在“课上听过”写代码时就会卡在“每个概率到底怎么算”上。第二拉普拉斯平滑是否理解。很多人知道公式是(count 1) / (N classes)但不知道分子分母为什么这么加。一旦测试样本出现训练集中没见过的特征值平滑就是唯一能保证预测不崩的手段。第三编码能力是否干净。你要在十几分钟内实现读入、统计、训练、预测、输出。常见问题包括Map用不熟练导致统计结构混乱、浮点数下溢出、读入顺序写错等。第四工程意识。数据结构怎么选、概率相乘要不要取对数、代码能不能直接提交到OJ这些都是“会做”和“能过”之间的差距。这道题和深度学习无关比的是基本功建议每一位算法岗候选人都能手写出来。2. 朴素贝叶斯原理与解题思路2.1 贝叶斯公式与大白话解释朴素贝叶斯的核心是贝叶斯公式P(C|X) P(X|C) * P(C) / P(X)其中C是类别0 或 1X (x1, x2, ..., xm)是特征向量。我们要做的是比较P(C0|X)和P(C1|X)哪个大哪个大就预测哪个。为什么叫“朴素”因为这里做了条件独立性假设在给定类别 C 的情况下每个特征之间相互独立。所以有P(X|C) P(x1|C) * P(x2|C) * ... * P(xm|C)这个假设在现实中往往不成立比如“天气晴”和“湿度低”其实有相关性但朴素贝叶斯依然能取得不错的效果而且实现极简。笔试场景下我们不需要讨论假设是否合理直接按公式实现即可。注意分母P(X)对所有类别都是同一个值所以比较后验概率大小时可以直接省略分母只比较分子P(C) * Π P(xi|C)这就是二分类器的决策规则。2.2 先验概率与条件概率怎么算先验概率P(C)用频率估计P(C0) count(C0) / n P(C1) count(C1) / n条件概率P(xi|C)用训练集中“类别为 C 且第 i 个特征取值为 xi”的样本数除以“类别为 C 的样本数”P(xi|C) count(C, xi) / count(C)举个例子训练集有 10 个样本其中 6 个标签为 0在某一个特征列上这 6 个样本中有 4 个取值为 3那么P(特征3 | C0) 4 / 6这就是“频率派”估计。不加平滑时如果测试样本某个特征取值在训练集该类别下从未出现那这个概率就会变成 0连乘后整个后验概率就变成 0模型直接输出错误结果。2.3 拉普拉斯平滑的必要性与公式拉普拉斯平滑的核心思想是为每个可能的取值加一个先验计数避免出现零概率。对于每个类别C下的第i个特征假设这个特征在类别C的训练样本中可能的取值集合大小为K_i(C)则平滑后的条件概率为P(xi|C) (count(C, xi) 1) / (count(C) K_i(C))分子加 1分母加的是该特征在该类别下的取值种类数。为什么分母要加K_i(C)因为我们要让这个类别下所有可能的取值概率之和等于 1。假设某特征在该类别下出现过 3 种取值的计数分别是 2, 3, 4平滑后分别变成 3, 4, 5总和为 12而分母是9 3 12正好归一。如果测试样本出现了该类别下从未见过的取值new_x此时count(C, new_x) 0平滑后概率为P(new_x|C) 1 / (count(C) K_i(C))不会为 0模型就可以继续计算。那么在笔试中我们怎么知道K_i(C)最简单的方式在训练阶段对每个类别下每个特征列维护一个 Set记录出现过哪些取值最后取 size。2.4 预测流程与防下溢出技巧预测时对每个测试样本执行计算P(C0)、P(C1)两个先验。对每个类别连乘所有P(xi|C)。比较两个连乘结果输出较大的类别。直接连乘有个问题当特征数 m 很大时每个概率都小于 1乘几十次后结果会小到超出浮点数的表示范围Java 的 double、C 的 double 都会变成 0。这叫做下溢出。解决办法是取对数。因为 log 是单调递增函数不影响大小比较log(P(C)) Σ log(P(xi|C))最后比较两个类别的 log 分数。注意拉普拉斯平滑后的概率不能取 0所以 log 不会出现负无穷。先验概率如果也做平滑可以写成log((count(C)1) / (n 2))因为只有两个类别所以分母加 2。当然如果不加先验平滑直接用频率也不会为 0但加上更稳。我在实际实现中会给先验也加平滑这样代码逻辑统一。3. 核心代码实现Java/C/Python3.1 公共思路与数据结构三种语言的逻辑完全一致我建议先梳理公用数据结构统计每个类别的样本数量classCount[label]统计每个类别下每个特征列的取值次数一个三维映射condCount[label][featureIndex][value]统计每个类别下每个特征列出现过多少种不同的取值valueSet[label][featureIndex]的 sizeJava 里可以用HashMapInteger, HashMapInteger, Integer嵌套表示condCount外层 key 是 label内层第一层 key 是特征下标内层第二层 key 是特征值。C 可以用mapint, mapint, mapint, int或者unordered_map。Python 则直接用数组套字典。这里有一个工程技巧由于标签只有 0 和 1可以用长度为 2 的数组classCount new int[2]条件概率统计也用Map[]数组下标 0 表示类别 0下标 1 表示类别 1。能省不少嵌套 Map 的写法。下面代码均默认训练集标签只有 0 和 1。3.2 Java 实现Java 是很多算法岗候选人提交时使用的语言注意提交时类名必须是Main不要粘贴包名。import java.util.*; public class Main { public static void main(String[] args) { Scanner sc new Scanner(System.in); int n sc.nextInt(); int m sc.nextInt(); int q sc.nextInt(); int[] classCount new int[2]; // condCount[c][i][val] 类别c下第i个特征取值为val的样本数 MapInteger, MapInteger, Integer[] condCount new Map[2]; // valueSet[c][i] 类别c下第i个特征出现过的不同取值 SetInteger[] valueSet new Set[2]; for (int c 0; c 2; c) { condCount[c] new HashMap(); valueSet[c] new HashSet(); } for (int row 0; row n; row) { int[] features new int[m]; for (int i 0; i m; i) { features[i] sc.nextInt(); } int label sc.nextInt(); classCount[label]; for (int i 0; i m; i) { int val features[i]; valueSet[label].add(val); // 这里valueSet是共用所有特征列的Set需要特别小心 MapInteger, Integer feaMap condCount[label].computeIfAbsent(i, k - new HashMap()); feaMap.put(val, feaMap.getOrDefault(val, 0) 1); } } // 由于上面的valueSet是单一Set不正确需要改成按特征区分 } }上面代码有个明显问题valueSet是单一 Set没有区分特征列。笔试时这种错误很伤正确写法是用SetInteger[]数组每个特征一个 Set。下面给出修正后的完整代码。修正后完整版import java.util.*; public class Main { public static void main(String[] args) { Scanner sc new Scanner(System.in); int n sc.nextInt(); int m sc.nextInt(); int q sc.nextInt(); int[] classCount new int[2]; MapInteger, MapInteger, Integer[] condCount new Map[2]; // valueSet[label][featureIndex]某个类别下某个特征出现过的取值集合 SetInteger[] valueSets new Set[2]; for (int c 0; c 2; c) { condCount[c] new HashMap(); valueSets[c] new HashSet(); } // 因为一个类别下每个特征列都要单独维护Set所以用数组嵌套会比较复杂。 // 直接用 MapInteger, MapInteger, SetInteger 更容易理解。 MapInteger, MapInteger, SetInteger featureValueSet new HashMap(); for (int label 0; label 2; label) { featureValueSet.put(label, new HashMap()); } for (int row 0; row n; row) { int[] features new int[m]; for (int i 0; i m; i) features[i] sc.nextInt(); int label sc.nextInt(); classCount[label]; MapInteger, SetInteger labelFeatureSet featureValueSet.get(label); for (int i 0; i m; i) { int val features[i]; // 维护 condCount MapInteger, Integer feaCountMap condCount[label].computeIfAbsent(i, k - new HashMap()); feaCountMap.put(val, feaCountMap.getOrDefault(val, 0) 1); // 维护取值种类 SetInteger set labelFeatureSet.computeIfAbsent(i, k - new HashSet()); set.add(val); } } for (int t 0; t q; t) { int[] test new int[m]; for (int i 0; i m; i) test[i] sc.nextInt(); double score0 Math.log((classCount[0] 1.0) / (n 2.0)); double score1 Math.log((classCount[1] 1.0) / (n 2.0)); MapInteger, SetInteger set0 featureValueSet.get(0); MapInteger, SetInteger set1 featureValueSet.get(1); for (int i 0; i m; i) { int val test[i]; MapInteger, Integer map0 condCount[0].get(i); int count0 map0 null ? 0 : map0.getOrDefault(val, 0); int k0 set0.containsKey(i) ? set0.get(i).size() : 0; // 平滑概率(count0 1) / (classCount[0] k0) double p0 (count0 1.0) / (classCount[0] k0); score0 Math.log(p0); MapInteger, Integer map1 condCount[1].get(i); int count1 map1 null ? 0 : map1.getOrDefault(val, 0); int k1 set1.containsKey(i) ? set1.get(i).size() : 0; double p1 (count1 1.0) / (classCount[1] k1); score1 Math.log(p1); } System.out.println(score0 score1 ? 0 : 1); } } }这里我刻意用了MapInteger, MapInteger, SetInteger来管理特征取值集合。注意当某个类别下某个特征列完全没有出现过任何值map0可能为 nullk0为 0此时表示训练集中该类别的样本数为 0但classCount[0]也可能为 0导致分母为 0。当然题目保证每个类别至少有一个样本吗不一定但为避免除零我通常会在读取后做检查或者直接把classCount初始分母改为Math.max(1, classCount[0])。不过如果某个类别完全没有样本分类器本身也没什么意义。建议在代码开头判断如果某个类别的样本数为 0就直接把所有测试样本预测为另一个类别。在面试题中通常不会出现这种极端数据这里只做提醒。另一个细节特征值范围高达10^9用int存储没有问题。如果题目改成字符串特征把Integer换成String即可。3.3 C 实现C 写这类题目最大的坑是数据结构嵌套复杂时容易写乱。我的建议是能不用unordered_map套unordered_map就不要用必要时可以直接用map对数规模数据完全没问题。下面给出一个可读性优先的版本。#include bits/stdc.h using namespace std; int main() { int n, m, q; cin n m q; // classCount[label] vectorint classCount(2, 0); // condCount[label][featureIndex][value] - count vectormapint, mapint, long long condCount(2); // featureValueSet[label][featureIndex] - set of values vectormapint, setint featureValueSet(2); for (int row 0; row n; row) { vectorint feats(m); for (int i 0; i m; i) cin feats[i]; int label; cin label; classCount[label]; for (int i 0; i m; i) { int val feats[i]; condCount[label][i][val]; featureValueSet[label][i].insert(val); } } cout fixed setprecision(10); for (int t 0; t q; t) { vectorint test(m); for (int i 0; i m; i) cin test[i]; double score0 log((classCount[0] 1.0) / (n 2.0)); double score1 log((classCount[1] 1.0) / (n 2.0)); for (int i 0; i m; i) { int val test[i]; // 类别 0 long long count0 condCount[0][i].count(val) ? condCount[0][i][val] : 0; int k0 featureValueSet[0][i].size(); double p0 (count0 1.0) / (classCount[0] k0); score0 log(p0); // 类别 1 long long count1 condCount[1][i].count(val) ? condCount[1][i][val] : 0; int k1 featureValueSet[1][i].size(); double p1 (count1 1.0) / (classCount[1] k1); score1 log(p1); } cout (score0 score1 ? 0 : 1) \n; } return 0; }这个代码在 C17 下直接可运行。注意两点condCount[0][i]如果之前没有对i建过 map用operator[]会自动创建一个空 map然后.count(val)可以安全调用。但为了严格保险你也可以先判condCount[0].count(i)。因为这里特征下标i是在循环中固定从 0 到 m-1 的训练时一定会对出现的特征列建过 map所以这里直接condCount[0][i]是安全的。如果训练集中某个特征列在某个类别下没有任何样本condCount[0][i]依然是存在于外层 map 中的因为外层 map 的 key 是特征下标只要类别 0 有样本且这些样本的特征列覆盖了所有 i就会建立。所以没问题。featureValueSet[0][i].size()同样安全因为训练循环里对每个特征列都执行了insert。3.4 Python 实现Python 的优势是写起来最短但要注意读入速度和浮点精度。建议使用sys.stdin.read()一次性读入所有数据然后用迭代器处理。完整代码如下import sys import math from collections import defaultdict def main(): data list(map(int, sys.stdin.read().split())) idx 0 n data[idx]; idx 1 m data[idx]; idx 1 q data[idx]; idx 1 class_count [0, 0] cond_count [defaultdict(lambda: defaultdict(int)) for _ in range(2)] value_set [defaultdict(set) for _ in range(2)] for _ in range(n): feats data[idx:idx m] idx m label data[idx]; idx 1 class_count[label] 1 for i, val in enumerate(feats): cond_count[label][i][val] 1 value_set[label][i].add(val) # 先验概率拉普拉斯平滑 log_prior0 math.log((class_count[0] 1) / (n 2)) log_prior1 math.log((class_count[1] 1) / (n 2)) out_lines [] for _ in range(q): test data[idx:idx m] idx m # 类别 0 的对数分数 score0 log_prior0 for i, val in enumerate(test): count0 cond_count[0][i].get(val, 0) k0 len(value_set[0][i]) p0 (count0 1) / (class_count[0] k0) score0 math.log(p0) # 类别 1 的对数分数 score1 log_prior1 for i, val in enumerate(test): count1 cond_count[1][i].get(val, 0) k1 len(value_set[1][i]) p1 (count1 1) / (class_count[1] k1) score1 math.log(p1) out_lines.append(0 if score0 score1 else 1) sys.stdout.write(\n.join(out_lines) \n) if __name__ __main__: main()这个 Python 版本用defaultdict省掉了大量“是否存在”的判断。需要注意cond_count[0][i].get(val, 0)中cond_count[0][i]会自动创建defaultdict(int)这是安全的但因为defaultdict的__getitem__会改变内部结构使用get并不会触发默认值创建所以不会污染统计结构。4. 实操中的常见问题与排查技巧4.1 未在训练集出现的特征取值导致概率为 0这是最常见的坑。很多同学不写拉普拉斯平滑直接count / classCount本地测试样例通过线上遇到一个“新特征值”就输出错误。排查方法很简单自己构造一个训练集中没有的取值作为测试数据观察程序是否崩溃或输出与预期不符。如果发现分子为 0基本就是平滑没写对。平滑时要特别注意分母里的K_i(C)。我看到过不少实现分子加 1分母只加 1比如(count 1) / (classCount 1)。这在单特征时没问题但如果某个特征在这个类别下有多个不同取值分母加 1 会导致该特征所有取值的概率之和不为 1。笔试数据小不容易验证但严格的验证方法是对某一个类别和特征列将所有可能取值的平滑概率求和看是否等于 1。例如counts {2, 3, 4} K 3 平滑概率 (21)/(93) (31)/(93) (41)/(93) 3/12 4/12 5/12 1如果分母加 1那就是 3/10 4/10 5/10 1.2显然这是错的。4.2 浮点下溢出与对数变换当m 20时如果每个概率都约 0.5连乘结果约0.5^20 9.5e-7还在 double 可表示范围内好像没问题。但当概率更小比如某些条件概率约 0.01 时0.01^20 1e-40依然可以表示。真正危险的是m很大或者概率非常小例如0.1^100 1e-100double 最小的正规格化数是2.2e-308所以 100 维时还没下溢出。但为了安全以及面试官可能会追问“为什么用对数”我强烈建议统一用对数实现。这样也和你手推公式时保持一致。还有一种情况是Math.log的参数为 0会导致负无穷。加了拉普拉斯平滑后任何概率都大于 0所以不会出现log(0)。如果你在调试中发现概率为 0先检查平滑是否生效。4.3 输入输出格式的细节三种语言的读入姿势不同容易踩的坑也不一样Java 的Scanner虽然方便但nextInt()不会处理行尾换行符这没问题。但如果你在第一个nextInt()前误用了nextLine()可能会读到空串。建议统一只用nextInt()。C 的cin x会自动跳过空白最稳妥。但如果使用scanf要小心%d和换行符的配合一般无需处理。Python 如果使用input().split()逐行读当数据行数多时会慢但n ≤ 1000完全没问题。我更推荐sys.stdin.read()一次性读入不容易因为末尾换行符导致解析错误。输出时注意每一行都要换行尤其 C 用\n而不是endl避免频繁 flush 降低性能。Python 用\n.join(...)也避免了逐行print带来的开销。4.4 代码提交时的几个致命错误在线笔试环境下Java 的类名必须是Main默认的public class Solution在某些 OJ 上会编译错误。C 提交时不要带#include bits/stdc.h这个大多数 OJ 支持但如果你不确定用标准的#include iostream、#include vector、#include map、#include set、#include cmath最保险。Python 则要注意不要提交 Jupyter notebook 格式也不要在文件中写交互代码。另外很多同学会在本地 IDE 加了package或import不存在的库提交前一定要注释掉。C 如果用了long long要确保读入时用cin val到long long变量类型不匹配会导致 UB。4.5 数据规模与时间复杂度的权衡n ≤ 1000m ≤ 20q ≤ 100。哪怕你用最朴素的遍历统计时间复杂度也完全够训练 O(nm)预测 O(qm)。但要注意如果特征值范围很大不能用数组直接落下标必须用哈希表或平衡树。这就是为什么代码里都用map/HashMap而不用固定大小数组。有些同学看到“非负整数”第一反应开一个int cnt[1000005]一旦取值超过这个范围就会越界。在笔试中一定要牢记10^9级别的值必须用哈希结构。5. 在线测试与环境准备5.1 本地自测方案从手搓数据到批量验证没有在线评测平台时建议按以下流程做本地自测准备一个input.txt内容格式如6 3 4 0 0 0 0 0 0 1 0 1 0 0 0 1 1 0 1 0 1 0 1 1 0 1 1 1 1 0 0 0 1 0 2 0 2 2 2运行程序读入input.txt输出结果。手工核算前几个样本。比如一个测试样本1 1 0训练集中类别 0 有 3 个样本类别 1 有 3 个样本先验相等。如果不考虑特征相关可以看到类别 1 中特征11的特征比较多预测为 1这符合直观。使用批量脚本对比三种语言的输出。比如在 Bash 中java Main input.txt java.out ./main input.txt cpp.out python3 main.py input.txt py.out diff java.out cpp.out diff cpp.out py.out三个输出一致基本能确认实现没有逻辑错误。这也是我在对比不同语言实现时最常用的方法。如果你想把题目挂到在线测试可以使用常见的在线评测系统OJ的“比赛模式”或“题目导入”也可以用 GitHub Actions/本地跑分脚本做自动化验证。重点在于输入格式、输出格式必须严格匹配题目描述尤其注意每行末尾是否允许多余空格。5.2 三种语言在 VSCode 下的环境配置要点很多候选人不是不会写是本地环境配不好导致调试效率极低。这里分享三个语言在 VSCode 下比较省心的配置Java安装 JDK17 或 JDK21然后在 VSCode 安装Extension Pack for Java。写好代码后直接用右上角运行按钮即可。注意Main.java文件名必须和类名一致。C安装 C 编译器。Windows 用户建议直接装 MSYS2/MinGW-w64 或 Visual Studio 的cl.exemacOS 则用clang。VSCode 安装C/C扩展后可以用tasks.json配置编译任务快捷键Cmd/Ctrl Shift B执行编译。如果遇到access violation c0000005这类运行时崩溃通常是指针越界或数据结构访问出错与编辑器环境无关优先检查代码逻辑。Python安装 Python 3然后在 VSCode 装 Python 扩展。先写脚本再用终端python3 main.py input.txt运行。这里提一句如果你需要安装第三方库如 sklearn推荐用pip install scikit-learn但这道题完全不需要第三方库。所谓“磨刀不误砍柴工”建议在笔试前把三种语言的最小运行模板准备好能读入整数、能循环处理、能格式化输出。这样遇到什么题都可以快速套用省去现场调试环境的时间。5.3 扩展从这道题到真实场景的朴素贝叶斯题目里的二分类器虽然简单但方法论可以直接迁移到文本分类、垃圾邮件识别、新闻分类等场景。比如“新闻分类”特征往往是词频或 TF-IDF 向量标签是多类别。朴素贝叶斯依然适用只不过样例中“特征列”变成了“特征词”而且特征维度可能成千上万。这时你更需要用对数概率和稀疏存储。我见过不少人先学了 sklearn 的MultinomialNB却不会手写结果笔试一碰到“实现朴素贝叶斯”就懵。建议在刷题时手动实现一遍这个基础模型能加深对条件独立性假设和平滑的理解。如果你用过 Python 的sklearn.naive_bayes.GaussianNB会发现它默认不用拉普拉斯平滑而是用高斯分布估计连续特征。这是朴素贝叶斯的另一种形态。在笔试中明确说了“离散特征”就用我们上面的多项分布模型不要混淆。6. 最后再分享一个小技巧我写这三种实现时其实是从 Python 版本先想清楚数据流再翻译成 Java 和 C 的。这样做的原因是 Python 表达逻辑最快写伪代码都不容易错但 Python 里defaultdict(lambda: defaultdict(int))这种嵌套结构翻译成 Java 时需要小心computeIfAbsent的用法翻译成 C 时则要预先想好 map 的层级。建议你先用 Python 跑通学习曲线再对照着写 Java/C比自己硬憋三种语言要快得多。还有一点这道题如果出现在笔试中建议先花 1 分钟在草稿纸上列出四个统计量——先验计数、条件计数、每类特征取值种类数、测试概率计算方式。把公式写在纸上再写代码能显著减少“写着写着忘了分母该加几”的情况。这个习惯让我在很多手写代码题里稳住了心态。最后不论你用什么语言一定要清楚朴素贝叶斯不是一个黑盒调包它背后就是“频率计数 平滑 对数连乘”。把这个核心抓住不管它伪装成什么题型你都能在面试现场快速写出干净的实现。
返回列表