简介:这是一份围绕《机器学习实战》英文版及中文PDF、Python源码和配套数据集的完整学习资料包,面向希望从理论走向动手实践的机器学习初学者和开发者。压缩包共79个文件,约32MB,其中74个py源码按书章节组织,涵盖K近邻、决策树、朴素贝叶斯、逻辑回归、支持向量机、AdaBoost、回归树等算法的可运行示例;另有2份PDF电子书、2份Markdown目录说明和1个数据集压缩包,便于对照阅读与本地调试。内容覆盖监督学习、无监督学习、特征降维及集成方法等核心主题,读者可依据章节代码逐步复现模型训练、评估与优化流程,也能利用数据集自行扩展实验。资源结构清晰、按Ch02至Ch09分章存放,已有2943人学习下载,适合系统入门机器学习并通过实操加深理解。
1. 为什么「机器学习实战」值得逐行复现:PDF 是入口,Python 才是重点
很多人下完《机器学习实战》的 PDF,第一件事是翻目录,第二件事是收藏,然后就没有然后了。这本书最反直觉的地方,恰恰在于它看起来像一本教材,实际上是本代码工作簿:不把配套的 Python 代码敲完,你很难真正理解 KNN 的距离投票、贝叶斯的 log 概率这些算法背后的手感。这份 PDF 配套的 Python 源码包我反复刷了三遍,环境调好、数据备齐、每一章的代码逐行跑通,从约会分类到垃圾邮件过滤都有完整可复现代码。它适合三类人:刚学完 Python 语法找不到项目的入门者,算法学过但没动手写过实现的学生,以及天天调 scikit-learn 却看不懂内部逻辑的从业者。
2. 先把 Python 环境调通:虚拟环境、依赖版本与三个翻车点
算法实战最花时间的往往不是算法本身,而是环境。这本书出版年代早,配套代码经历过 Python 2 到 Python 3 的大迁移,如果你一上来就用最新解释器跑,头几个报错就能劝退一半人。先把环境固定好,后面每一章才能复现出稳定结果。
2.1 为什么固定 Python 3.8:老代码的字符串与迭代器兼容边界
这本书最早的示例代码大量使用 Python 2 语法,后来社区迁移到 Python 3,但很多下载包里迁移得不彻底。最典型的三个差异:dict 的iteritems()在 Python 3 中改名items(),zip()从返回列表变成返回迭代器,map()的结果也不再是 list。这些差异在书里朴素贝叶斯和 SVM 部分尤其常见。如果直接用 Python 3.12 跑,遇到AttributeError的概率接近百分之百。
最省心的土办法是固定一个兼容版本:Python 3.8.10。它既能运行旧式字符串处理习惯,又对现代 numpy、scikit-learn 有较好支持。用 conda 创建独立环境,命令如下:
conda create -n mla python=3.8.10 -y conda activate mla python --version这里-n mla是环境名,python=3.8.10指定解释器小版本,-y跳过确认步骤。激活后python --version应输出 3.8.10。 如果不习惯 conda,用python -m venv mla也可以,只是 conda 在切换不同 Python 版本时更直观。需要注意:如果你激活了环境,但 VSCode 里仍选中 base 解释器,实际跑的并不是 3.8.10,所以第一步确认版本远比跑代码重要。
这个版本选择不是玄学,而是迁移边界。numpy 1.24.4 是最后一个支持 Python 3.8 的 1.x 版本,scipy 1.10.1 也在这个范围内;到了 Python 3.12,numpy 需要升到 1.26,但老代码里很多np.float、np.int的写法早就被移除,反而引入新问题。固定版本本质上是给老代码一个「后悔药」,先能跑通,再谈优化。
2.2 依赖版本锁定:numpy、scipy、matplotlib 的三角关系
这本书的代码不依赖重量级框架,最核心的就是 numpy、scipy、matplotlib。加上 scikit-learn 是为了后期对照准确率,并不是必须。建议在项目根目录建一个 requirements.txt,直接锁定版本:
numpy==1.24.4 scipy==1.10.1 matplotlib==3.7.5 scikit-learn==1.3.2安装时用国内镜像会快很多:
pip install -r requirements.txt -i https://pypi.tuna.tsinghua.edu.cn/simple-i指定 pip 镜像源,清华源对大多数包都有缓存。为什么这样锁:numpy 1.24.4 与 scipy 1.10.1 是官方测试过的组合,不会出现numpy.dtype size changed这类 ABI mismatch;matplotlib 3.7.5 能正常渲染书里散点图,又不会像 3.8 那样和 numpy 争抢编译头文件。scikit-learn 1.3.2 支持 Python 3.8,且KNeighborsClassifier的接口和书里手写版本接近,方便后期做交叉验证。
如果你在安装时遇到Could not find a version that satisfies the requirement scipy==1.10.1,大概率是 Python 版本不匹配。可以先跑python --version确认环境,再检查是否在 mla 环境内。另一个常见问题是 pip 缓存了旧轮子,建议加一句pip install --upgrade pip后再装。
2.3 VSCode 下的 Python 解释器选择与 launch.json 配置
VSCode 是最适合逐行复现这本书的编辑器,关键是把解释器指到刚才创建的环境。装了 Python 扩展后,按Ctrl+Shift+P调出命令面板,输入Python: Select Interpreter,再选中mla环境。这一步直接决定右上角运行按钮用的是哪个 Python,所以我一般会在项目根目录放一个.vscode/settings.json,内容类似:
{ "python.defaultInterpreterPath": "/opt/anaconda3/envs/mla/bin/python", "files.encoding": "utf8" }Windows 下路径通常是C:\\Users\\你的用户名\\anaconda3\\envs\\mla\\python.exe。files.encoding设为 utf8 是为了减少读取数据文件时的编码干扰。如果你打开别人的代码出现中文注释乱码,这个设置也能缓解。
调试时还需要一份 launch.json:
{ "name": "Python: 当前文件", "type": "python", "request": "launch", "program": "${file}", "console": "integratedTerminal", "justMyCode": false, "cwd": "${workspaceFolder}" }cwd设为工作区根目录,能保证相对路径按项目根目录解析;justMyCode设为 false,调试时可以进入第三方库内部,排查 numpy 数组形状问题很管用。console用集成终端,print输出不会被截断。如果你按 F5 后发现第一行import numpy就报错,先看左下角解释器,换成 mla 再重试。
2.4 最小烟雾测试:两行代码验证环境通没通
环境配好不等于能跑,我习惯先做一个最短的烟雾测试,把 numpy、matplotlib 一起验掉:
import sys import numpy as np import matplotlib.pyplot as plt print(sys.version) print(np.__version__, plt.__version__) x = np.linspace(0, 2 * np.pi, 100) plt.plot(x, np.sin(x)) plt.savefig("smoke_test.png") print("env ok")这段代码只做三件事:打印解释器版本、确认两个库能 import、生成一张正弦图。如果smoke_test.png成功出现在目录里,说明环境基本可用。np.linspace(0, 2*np.pi, 100)生成 0 到 2π 的 100 个点,plt.savefig不弹窗,适合在无桌面环境或容器里验证。若报No module named matplotlib,就回到 2.2 重新安装;若报ImportError: DLL load failed,优先检查 numpy 与 scipy 版本搭配,因为 Windows 下 DLL 冲突比 Linux 更常见。
这套流程跑通后,再打开书里的第一个算法,基本不会因为环境问题打断复现节奏。
3. 复现 kNN:从约会分类到手写数字识别
环境通后,最值得先跑的是书第 2 章的 kNN。它没有训练阶段,概念最简单,但能让你同时看到数据的读取、归一化、距离计算和投票机制。先跑通它,后面朴素贝叶斯和树模型都会顺很多。
3.1 kNN 核心逻辑:距离计算、投票与归一化
kNN 的「训练」其实就是保存整个训练集,预测时对新样本计算它与每个训练样本的欧氏距离,取距离最近的 k 个邻居,让邻居投票决定类别。这个逻辑在原书里浓缩成一个 classify0 函数:
import numpy as np def classify0(inX, dataSet, labels, k): dataSetSize = dataSet.shape[0] diffMat = np.tile(inX, (dataSetSize, 1)) - dataSet sqDistances = (diffMat ** 2).sum(axis=1) distances = np.sqrt(sqDistances) sortedDistIndicies = distances.argsort() classCount = {} for i in range(k): voteIlabel = labels[sortedDistIndicies[i]] classCount[voteIlabel] = classCount.get(voteIlabel, 0) + 1 sortedClassCount = sorted(classCount.items(), key=lambda x: x[1], reverse=True) return sortedClassCount[0][0]np.tile(inX, (dataSetSize, 1))把待预测样本复制成与训练集同样行数的矩阵,减法是逐元素做的;.sum(axis=1)按行求和,得到每个训练样本的距离平方;argsort()返回从小到大排序后的索引,而不是排序后的值,这样可以用索引回取 labels。classCount.get(voteIlabel, 0) + 1是标准计数写法,第一次遇到某类别时默认 0,随后累加。
k 值一般取奇数,书中约会分类和手写数字都用 k=3。k 太小,模型对单个噪声点敏感;k 太大,人多的类会压过局部特征。你在复现时可以试着把 k 改成 1、5、15,观察错误率变化,这个实验比背公式更能建立手感。
3.2 文件转矩阵与归一化:数据是 kNN 的真瓶颈
kNN 的计算瓶颈在距离矩阵,而距离矩阵的质量取决于特征尺度。书中约会网站数据集每行四列,前三维特征分别是「每年飞行里程」「玩游戏所占时间百分比」「每周消费冰淇淋公升数」。这三个特征数量级完全不同,不归一化的话,飞行里程会完全主导距离。第一步是把 txt 读成矩阵:
def file2matrix(filename): with open(filename, encoding='utf-8') as fr: lines = fr.readlines() number = len(lines) returnMat = np.zeros((number, 3)) classLabelVector = [] for index, line in enumerate(lines): listFromLine = line.strip().split('\t') returnMat[index, :] = listFromLine[0:3] classLabelVector.append(int(listFromLine[-1])) return returnMat, classLabelVectorwith open(...)保证文件用完即关;encoding='utf-8'解决了 Windows 下默认 gbk 解码报错的问题;strip()去掉换行,split('\t')按 Tab 拆分。listFromLine[0:3]取特征,int(listFromLine[-1])取标签。如果你的数据是 csv,把'\t'换成','即可,其余不用改。
归一化函数是另一个不能跳过的环节:
def autoNorm(dataSet): minVals = dataSet.min(0) maxVals = dataSet.max(0) ranges = maxVals - minVals m = dataSet.shape[0] normDataSet = (dataSet - np.tile(minVals, (m, 1))) / np.tile(ranges, (m, 1)) return normDataSet, ranges, minValsmin(0)按列取最小,等价于每个特征的最小值;ranges是每列跨度。np.tile(minVals, (m, 1))把最小值向量复制成 m 行,每一行样本都减去自己那列的最小值,再除以跨度,最后所有特征落在 0 到 1 之间。边界情况是某一列所有值相同,ranges为 0,除零后得到 nan。我一般会在归一化前加一句dataSet[:, ranges == 0] = 0或直接跳过这一列,避免后面计算距离出现 nan 还找不到原因。
3.3 手写数字识别:把 32×32 的图片文件压成 1024 维向量
手写数字识别是 kNN 章节里更完整的例子。数据集里每个数字是一个 32×32 的文本文件,比如0_0.txt,文件名第一部分就是真实标签。读取时必须逐行读并展开成一行 1024 维向量:
def img2vector(filename): returnVect = np.zeros((1, 1024)) with open(filename, encoding='utf-8') as fr: for i in range(32): lineStr = fr.readline() for j in range(32): returnVect[0, 32 * i + j] = int(lineStr[j]) return returnVect每个文件恰好 32 行,每行 32 个 0/1 字符。外层for i遍历行,内层for j遍历列,32 * i + j把二维坐标映射到一维索引。这里逐行readline()比readlines()更省内存,因为每个文件只有 32 行,但测试集几百个文件时习惯养成后不容易出问题。
完整评估逻辑如下:
import os def handwritingClassTest(): train_dir = 'trainingDigits' test_dir = 'testDigits' trainList = os.listdir(train_dir) m = len(trainList) trainMat = np.zeros((m, 1024)) hwLabels = [] for i in range(m): fileNameStr = trainList[i] hwLabels.append(int(fileNameStr.split('_')[0])) trainMat[i, :] = img2vector(os.path.join(train_dir, fileNameStr)) testList = os.listdir(test_dir) errorCount = 0.0 mTest = len(testList) for i in range(mTest): fileNameStr = testList[i] truth = int(fileNameStr.split('_')[0]) testVect = img2vector(os.path.join(test_dir, fileNameStr)) predict = classify0(testVect, trainMat, hwLabels, 3) if predict != truth: errorCount += 1 return errorCount / mTestos.path.join负责拼路径,兼容 Windows 和 Linux。os.listdir返回的文件顺序不保证按名称排序,但这里每个样本都是独立读取,顺序不影响结果。classify0的 k 设为 3,在这个数据集上错误率大约 1.2%;如果你发现错误率超过 5%,先检查训练集和测试集是否放反,再检查img2vector是否正确把行索引映射到了 1024 维位置。
这一章的完整源码包里已经包含datingTestSet2.txt、trainingDigits、testDigits三份数据,直接把上面三个函数串起来就能得到最终错误率。先跑通 kNN,你会对「特征表达决定模型上限」这句话有很直观的感受。
4. 朴素贝叶斯:从垃圾邮件分类看文本特征工程
第四章的朴素贝叶斯是很多入门者第一次接触文本数据的地方。和 kNN 不同,这里要处理的是不定长的邮件内容,核心工作从距离计算变成了词表构建和文本清洗。跑完这一章,你会明白为什么说特征工程有时比算法更重要。
4.1 条件独立假设:为什么朴素还能用
朴素贝叶斯计算的是后验概率 P(类别|词向量),公式里假设词与词之间条件独立。邮件里「免费」和「点击」明显不独立,但假设独立后只需要统计每个词在每个类别下的频率,计算量大幅下降,实际效果也不错。这就是「朴素」二字的由来。
代码层面要做三件事:给每封邮件生成词向量、统计每个类别的词频、转成对数概率。训练部分的核心函数如下:
def trainNB0(trainMatrix, trainCategory): numTrainDocs = len(trainMatrix) numWords = len(trainMatrix[0]) pAbusive = sum(trainCategory) / float(numTrainDocs) p0Num = np.ones(numWords) p1Num = np.ones(numWords) p0Denom = 2.0 p1Denom = 2.0 for i in range(numTrainDocs): if trainCategory[i] == 1: p1Num += trainMatrix[i] p1Denom += sum(trainMatrix[i]) else: p0Num += trainMatrix[i] p0Denom += sum(trainMatrix[i]) p1Vect = np.log(p1Num / p1Denom) p0Vect = np.log(p0Num / p0Denom) return p0Vect, p1Vect, pAbusivepAbusive是垃圾邮件先验概率。p0Num、p1Num用 ones 初始化,分母用 2.0,这是拉普拉斯平滑,防止某个词在测试集出现但训练集没出现时概率直接变成 0。np.log把所有概率转成对数,后面分类时连乘变连加,避免浮点数下溢到 0。为什么不用概率相乘?几十个词叠下去,概率乘积很快就低于1e-300,在内存里就是 0.0,取 log 后数值范围稳定得多。
分类函数比较两份对数后验:
def classifyNB(vec2Classify, p0Vec, p1Vec, pClass1): p1 = sum(vec2Classify * p1Vec) + np.log(pClass1) p0 = sum(vec2Classify * p0Vec) + np.log(1 - pClass1) return 1 if p1 > p0 else 0vec2Classify是测试邮件的词向量,p1Vec已经是每个词在垃圾类别下的对数条件概率。vec2Classify * p1Vec保留出现词的 log 概率,没出现的词乘 0 自动忽略,最后加上 log 先验。如果你希望宁可漏接也不误杀,可以改成p1 > p0 + 0.3,这个 0.3 是置信度偏移,相当于人为提高判断为垃圾的门槛。
4.2 文本切分与词表构建:不能只会按空格 split
书里的底线例子用split()切词,对一句话还行,对真实邮件完全不够用。邮件里有 HTML 标签、标点、大小写和 URL,直接按空格切会把hello!和hello当成两个词。常见的做法是用正则把所有非单词字符当分隔符:
import re def textParse(bigString): listOfTokens = re.split(r'\W+', bigString) return [tok.lower() for tok in listOfTokens if len(tok) > 2] def createVocabList(dataSet): vocabSet = set() for document in dataSet: vocabSet |= set(document) return list(vocabSet) def setOfWords2Vec(vocabList, inputSet): returnVec = [0] * len(vocabList) for word in inputSet: if word in vocabList: returnVec[vocabList.index(word)] = 1 return returnVec\W+匹配连续的非字母数字字符,一次拆掉空格、逗号、句号、HTML 标签符号。len(tok) > 2过滤掉 I、a、to 这类高频但判别力弱的短词。注意list(vocabSet)的顺序不固定,但 0/1 向量只要每次构建词表后保持同一顺序,训练和测试用同一个词表,就不会影响分类。性能方面,vocabList.index(word)是 O(N) 查找,词表几千词时还好,如果换成几万词,建议先用{word: i for i, word in enumerate(vocabList)}建索引再查。
词集模型只记录某个词是否出现,词袋模型则记录出现次数:
def bagOfWords2VecMN(vocabList, inputSet): returnVec = [0] * len(vocabList) for word in inputSet: if word in vocabList: returnVec[vocabList.index(word)] += 1 return returnVec「免费」出现 5 次的邮件比出现 1 次更像垃圾邮件,所以词袋在垃圾邮件场景下通常更有效。原书主力代码用的是词集模型,你可以在复现时切换成词袋版本重新跑一次,对比错误率。这个对比实验能让你直观看到特征表达粒度对模型的影响。
4.3 交叉验证:别用同一批数据训练和测试
书里 spamTest 的思路是随机抽 10 封邮件做测试,其余训练,重复统计错误率。原书有一段代码用random.uniform生成随机下标,存在越界风险。更稳的写法是先打乱 50 个样本下标,再切出训练集和测试集:
import random def spamTest(): docList = [] classList = [] for i in range(1, 26): with open(f'email/spam/{i}.txt', encoding='utf-8') as f: docList.append(textParse(f.read())) classList.append(1) with open(f'email/ham/{i}.txt', encoding='utf-8') as f: docList.append(textParse(f.read())) classList.append(0) vocabList = createVocabList(docList) idx = list(range(50)) random.shuffle(idx) testSet = idx[:10] trainSet = idx[10:] trainMat = [] trainClasses = [] for i in trainSet: trainMat.append(setOfWords2Vec(vocabList, docList[i])) trainClasses.append(classList[i]) p0V, p1V, pSpam = trainNB0(np.array(trainMat), np.array(trainClasses)) errorCount = 0 for i in testSet: wordVector = setOfWords2Vec(vocabList, docList[i]) if classifyNB(np.array(wordVector), p0V, p1V, pSpam) != classList[i]: errorCount += 1 return errorCount / len(testSet)random.shuffle(idx)是在原列表上打乱,前 10 个做测试,后 40 个训练。每次运行错误率会有波动,通常在 10% 以内。想要可复现,可以加random.seed(42)。np.array(trainMat)这一步需要确认 trainMat 是纯 0/1 整数列表,否则 dtype 变成字符串后,trainNB0 里的数值加和会变成字符串拼接,最后算出一堆 nan。
我在复现时发现,影响错误率最大的不是贝叶斯公式,而是textParse里len(tok) > 2这个过滤条件。去掉它,错误率会明显上升,因为短词在垃圾和正常邮件里分布都很均匀,相当于给分类器加了大量噪声。这个点正好说明:项目里最值得调试的不是算法参数,而是数据进入模型前那一段清洗逻辑。
5. 避坑/常见问题:运行老书代码的 5 个典型坑
这本书的配套代码在网上流传了很多版本,有原版 Python 2 的,有半迁移的,还有改了路径的。无论你下载的是哪份,运行到固定章节总会撞上下面几个坑。我按现象、原因、解决三个步骤记录,照着排查就行。
5.1 AttributeError: 'dict' object has no attribute 'iteritems'
现象:kNN 或朴素贝叶斯的统计代码里,for key, val in classCount.iteritems():直接报 AttributeError,整个脚本停掉。
原因:Python 2 的字典迭代方法iteritems()在 Python 3 中被彻底移除,字典遍历统一用items()。老代码没有迁移干净,就会在这里翻车。
解决:全局搜索iteritems,替换成items。需要注意in后面的变量不要弄错,比如classCount.items()返回的视图对象在遍历时不能新增键;但这里的classCount是已经统计完的字典,只读遍历没问题。修改后sorted(classCount.items(), key=lambda x: x[1], reverse=True)就是标准的按票数降序排序写法。这是全书出现频率最高的兼容性问题,先替换再跑能省很多时间。
5.2 UnicodeDecodeError: 'gbk' codec can't decode byte
现象:Windows 下直接open('email/spam/1.txt')读邮件数据,抛 UnicodeDecodeError,提示 gbk 解码失败。
原因:Python 3 的open()默认编码跟随操作系统 locale,中文 Windows 通常是 gbk。而资源包里的数据文件多数是 utf-8 或纯 ASCII,用 gbk 可能解不开某些字符,尤其当文件里混入特殊符号时。
解决:所有读取数据的地方显式指定编码:
with open(filename, encoding='utf-8') as f: text = f.read()如果还不确定文件编码,可以加errors='ignore'快速跳过坏字节,但正式复现时不要依赖它,因为丢字符会影响词表构建。我的习惯是先把数据文件用编辑器另存为 utf-8,再跑代码,这样能从源头消掉问题。手写数字那个数据集虽然是纯 0/1 文本,但用encode='utf-8'打开也不会出问题。
5.3 TypeError: only size-1 arrays can be converted to Python scalars
现象:运行朴素贝叶斯时,trainNB0里p1Num += trainMatrix[i]报类型错误,或者训练完成后 p1Vect 全是 nan。
原因:trainMatrix是 list 的 list,里面混入字符串、None 或布尔值时,np.array(trainMatrix)会把整个数组推断成字符串类型。后续数值加法变成字符串拼接,最终概率全是 nan。最常见的触发点是文本解析时没过滤干净,词表里混进None,或者数据文件开头有 BOM 头。
解决:在进入trainNB0之前强制转换:
trainMat = np.array(trainMat, dtype=np.float64)同时打印一下trainMat.dtype,如果是object或<U...,说明有非数值成分。回到textParse检查分词结果,用all(isinstance(w, str) for w in tokens)快速验证。这个坑最容易出现在从网上直接拷贝的迁移版代码里,因为原作者可能只在某个特定数据上测试过。
5.4 matplotlib 中文乱码与图像窗口一闪而过
现象:plt.title('约会数据集')画图后,标题是方块,或者plt.show()不弹窗直接结束。
原因:matplotlib 默认字体里没有中文字形,Windows 下中文渲染会变成方块。而show()不弹窗通常是因为当前环境没有可用的 GUI 后端,比如在无桌面 Linux 或容器里运行。
解决:在绘图脚本顶部加两行:
plt.rcParams['font.sans-serif'] = ['SimHei'] plt.rcParams['axes.unicode_minus'] = FalseSimHei是 Windows 常见中文字体,Linux 下可以改成Noto Sans CJK SC。如果还是没有窗口,先改成plt.savefig('plot.png'),确认能出图后再排查 GUI 后端。VSCode 里如果装了 Jupyter 插件,也可以把plt.show()替换成%matplotlib inline直接在面板里看图。我在复现时习惯先把图表标题写成英文,跑通流程后再改中文,省得在字体问题上花半小时。
5.5 FileNotFoundError: trainingDigits 目录找不到
现象:从资源包里拷贝单独的.py文件到新目录运行,报No such file or directory: 'trainingDigits'。
原因:原书代码使用相对路径,而相对路径是相对于「当前工作目录」解析的。在 VSCode 里直接打开代码所在目录没问题;但如果你从项目的子目录启动脚本,或者用命令行python code/knn.py运行,工作目录就变成了项目根目录,原始相对的trainingDigits自然找不到。
解决:改成基于脚本文件所在目录拼接路径:
from pathlib import Path BASE_DIR = Path(__file__).resolve().parent train_dir = BASE_DIR / 'trainingDigits' test_dir = BASE_DIR / 'testDigits'__file__是当前脚本完整路径,resolve().parent取它所在目录。这样无论从哪里启动,都能找到数据。老式写法os.path.join(os.path.dirname(os.path.abspath(__file__)), 'trainingDigits')也可以,但 pathlib 的/拼接在 Windows 上更简洁。我每复现一章都会把资源包整理成code/与data/两层,然后让代码用BASE_DIR.parent / 'data'定位数据,基本不会再遇到路径问题。
6. 进阶:把书里的例子改装成自己的小项目
6.1 给 kNN 加一个返回置信度的接口,而不是只返回类别
书里的classify0只返回最终类别,你在调参时只能看到「对或错」,看不到模型有多确信。我把它改装成同时返回得票比例,这样能直接看出 k 值变化带来的影响:
def knn_with_proba(inX, dataSet, labels, k): dataSetSize = dataSet.shape[0] diffMat = np.tile(inX, (dataSetSize, 1)) - dataSet distances = np.sqrt((diffMat ** 2).sum(axis=1)) sortedIdx = distances.argsort() classCount = {} for i in range(k): label = labels[sortedIdx[i]] classCount[label] = classCount.get(label, 0) + 1 votes = sorted(classCount.items(), key=lambda x: x[1], reverse=True) prob = votes[0][1] / k return votes[0][0], prob返回的prob是最大得票数除以 k,范围在 1/k 到 1 之间。在约会数据上,你可以循环 k=1 到 10,画一条错误率随 k 变化的曲线:k=3 到 5 通常最低,k 过大错误率开始上升。这个过程比看任何公式都直观。
朴素贝叶斯同样可以输出自信度:把classifyNB里的p1和p0差值取绝对值,差值越大说明模型越确信。我在跑邮件分类时会把置信度低于 0.05 的测试样本单独打印出来,逐个检查是词表缺词还是文本清洗不干净,往往能发现比调阈值更值得优化的地方。
做完这两步,你已经不是「照着书跑代码」,而是把书里的算法封装成自己能改的工具。后续再做小的分类项目时,比如给自己的 csv 数据做分类,直接复用file2matrix和autoNorm,注意排除日期、ID 这类无关特征,再套上knn_with_proba就能快速看到基线结果。从那以后,我每次拿到新书里的代码,都会先花二十分钟把环境固定、路径固定,再开始复现,省下的时间往往不止二十分钟。希望帮到你。
本文还有配套的精品资源,点击获取