news 2026/9/26 18:05:18

八种机器学习算法在MNIST手写数字识别中的实践与避坑指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
八种机器学习算法在MNIST手写数字识别中的实践与避坑指南

简介:面向机器学习初学者与算法实践者,这份压缩包将八种经典分类算法——AdaBoost、朴素贝叶斯、决策树、KNN、逻辑斯蒂回归、最大熵、SVM与感知机——统一用于MNIST手写数字识别任务,提供一套完整的Python参考实现与使用案例。包内共有15个文件,其中13个为.py脚本,另附1个说明文档和1个辅助文本,整体仅27KB,轻量便携,适合快速下载与本地运行。各类算法均按独立目录组织,各自配有可执行的主程序,部分还包含独立子模块;同时提供决策树相关辅助文件与README说明,便于按步骤重现结果、调整超参数并对比不同模型在同一数据上的分类精度与收敛表现。对希望掌握多种分类器原理、完成机器学习课程作业、或进行算法横向对比的读者,这套代码既可作为入门模板,也能作为扩展开发的起点。目前已有726人学习下载,兼具实用性与参考价值。

1. 一份跑通八种MNIST算法的源码包:先看它能替你省下什么

手头拿到一个sklearn-MNIST-main的压缩包,打开一看,adaboost、贝叶斯朴素法、决策树、KNN、逻辑斯蒂、最大熵、SVM、感知机八个算法目录排得整整齐齐,每个目录下都挂着main.py或main1.py。这不是论文代码,也不是教学演示,而是一份能直接跑通的手写数字识别底稿——它把MNIST数据装载、训练、预测、评估这条链路在每个算法里都走了一遍。对正在做课程设计、入门机器学习对比实验、或者想快速抄一份能交差代码的人来说,最值钱的不是某个算法的精度,而是八个算法共用同一套数据管道的写法:换算法只换模型行,前后处理完全复用。我拆完这份包之后,把每个脚本从头到尾过了一遍,下面这篇笔记就是按“数据怎么进来、八个算法各自怎么落地、坑在哪、跑通后怎么变成工具”的顺序写的。

2. 先把数据弄明白:MNIST怎么装、装成什么形状、这份包怎么组织

2.1 数据装载与标签转换:npz格式和keras接口的取舍

MNIST 最常见的获取方式有两种:一是keras.datasets.mnist.load_data(),二是直接读mnist.npz文件。很多初学者在这第一步就翻车,因为keras接口在不同版本下的返回格式略有差异,而且它依赖后端环境,TensorFlow 没配好就报AttributeError。这份源码包里用的是后者——从data目录直接读mnist.npz,再用 numpy 手动拆开,好处是不依赖深度学习框架,sklearn 环境下就能跑。

import numpy as np # 从本地读取 mnist.npz,对应包里的 data 目录 with np.load('data/mnist.npz') as f: x_train, y_train = f['x_train'], f['y_train'] x_test, y_test = f['x_test'], f['y_test'] # 原始数据是 28x28 的二维矩阵,分类器需要一维特征,这里直接拉平 x_train = x_train.reshape(x_train.shape[0], 784).astype('float32') x_test = x_test.reshape(x_test.shape[0], 784).astype('float32') # 标签统一转成 int,后面做混淆矩阵和分类报告时才不会踩类型的坑 y_train = y_train.astype(int) y_test = y_test.astype(int) print(f'训练集: {x_train.shape}, 测试集: {x_test.shape}') print(f'标签范围: {y_train.min()} - {y_train.max()}')

这段代码的核心是reshape(x_train.shape[0], 784):MNIST 每张图是 28 乘 28 像素,展开后就是 784 维向量,这个维度对后面所有算法都一样。astype('float32')是为了减少内存占用,八种算法里 KNN 和 SVM 特别吃内存,如果默认读进来是float64,训练集 60000 乘 784 的矩阵直接占掉几百 MB,再加上算法内部的复制,很容易把 8GB 内存的机器跑崩。标签转int是很多脚本里容易漏掉的一步——sklearn 的accuracy_score在标签是uint8时也能算,但一旦做classification_report输出类别名或者画混淆矩阵,类型不统一就会报警告。

2.2 一份能直接跑的sklearn-MNIST目录导览:数据目录和各算法入口别走错

拆开压缩包后,目录结构很清楚,但第一次打开的人容易犯一个迷糊:main.py和main1.py到底有什么区别。我在decision_tree、svm、max_shang这些目录里看到两个入口文件,实际对一遍代码发现,main.py是完整流程版——装载数据、切分、训练、评估一步不落;main1.py是精简版,有的只输出准确率,有的把模型参数写死成了硬编码。跑的时候用哪个都行,建议以main.py为准,因为它的输出信息更全,能看到分类报告和混淆矩阵。

这份包里还有一个值得注意的点:max_Ent.py放在max_shang目录下,但目录里同时存在main.py,两者不是同一份代码。拆开看,max_Ent.py是用“最大熵模型的迭代求解”思路写的,本质上是把逻辑斯蒂回归的多分类版本用梯度下降手工实现了一遍;main.py则是直接调 sklearn 的LogisticRegression做对照。这些入口文件的位置是作者按自己习惯放的,你复制到别的项目里时,最好统一重命名,避免后面维护的时候不知道跑哪个。

2.3 train_test_split与归一化:为什么MNIST必须除255而不是做标准化

from sklearn.model_selection import train_test_split # 进一步切出验证集,方便调参阶段快速看效果 x_train_sub, x_val, y_train_sub, y_val = train_test_split( x_train, y_train, test_size=0.2, random_state=42, stratify=y_train ) # 像素归一化到 [0, 1],这是 MNIST 场景下最稳的预处理 x_train_sub = x_train_sub / 255.0 x_val = x_val / 255.0 x_test = x_test / 255.0 print(f'验证集大小: {x_val.shape[0]}')

这里用test_size=0.2表示从 60000 张训练图里留出 12000 张做验证,random_state=42保证每次切分的结果一致,stratify=y_train让切分后的类别比例和原始数据集一致——MNIST 里数字 0 和 1 的样本量本身就有差异,不 stratify 的话,某些数字可能在验证集里明显变少,影响模型对比的可信度。归一化除 255 而不是做StandardScaler标准化,是因为图片像素本身就是 0 到 255 的亮度值,除 255 后落在 0 到 1 区间,物理意义保留,也避免标准化后负数像素把某些分类器的决策边界搞偏。SVM 和 KNN 对特征尺度极敏感,这一步不做,后面说啥都白搭。

3. 八个算法逐个落地:每个脚本的入口行、参数和动手改的位置

3.1 KNN:最容易出效果但也最吃内存的基线

KNN 在这份包里被放在knn目录,main.py的实现用的是最经典的KNeighborsClassifier。MNIST 用 KNN 的直觉很简单:784 维空间里,同数字的图片距离近,不同数字的距离远。它不需要训练过程,但预测时要拿新样本和所有训练样本算距离,所以内存和耗时都在这里。

from sklearn.neighbors import KNeighborsClassifier # n_neighbors=5 是 sklearn 默认值,MNIST 上可以先从 3 试起 knn = KNeighborsClassifier( n_neighbors=5, weights='distance', # 距离加权:近的邻居投票权重大 algorithm='kd_tree', # 数据维度高,kd_tree 比 brute 省时间 n_jobs=-1 # 用满所有 CPU 核,KNN 是天然可并行的 ) knn.fit(x_train_sub, y_train_sub) val_acc = knn.score(x_val, y_val) print(f'KNN 验证集准确率: {val_acc:.4f}')

weights='distance'是 KNN 在 MNIST 上是否好用的关键开关:默认的uniform让所有邻居投票权重一样,但如果某几个邻居离得非常远,它们的票和最近邻同权,很容易带偏。改成distance后,距离近的样本话语权更大,能明显提升准确率。algorithm='kd_tree'在高维数据上有争议,因为 784 维下 kd_tree 的切分效率会退化,但实际跑起来比暴力brute略快,如果你要追求极致速度,可以改成brute然后靠n_jobs=-1拉满并行。

3.2 逻辑斯蒂回归:max_iter和solver参数别抄默认值

logistics/logistics.py这份代码和别的目录有个明显区别:它不是直接调 sklearn,而是先写了一个二元逻辑斯蒂的梯度下降实现,再在main.py里调LogisticRegression做多分类对照。这部分代码的价值在于让你看到手写实现和库实现的差距——手写版本跑得慢还不一定收敛,库版本收敛稳定且接口完整。

from sklearn.linear_model import LogisticRegression # MNIST 是多分类,solver 用 lbfgs 或 newton-cg 都行,别用 liblinear log_reg = LogisticRegression( solver='lbfgs', max_iter=300, # 默认 100 在 MNIST 上经常没收敛就停了 multi_class='multinomial', C=1.0 ) log_reg.fit(x_train_sub, y_train_sub) train_acc = log_reg.score(x_train_sub, y_train_sub) val_acc = log_reg.score(x_val, y_val) print(f'逻辑斯蒂 训练集准确率: {train_acc:.4f}, 验证集准确率: {val_acc:.4f}')

max_iter=300是这份代码里最实用的一行:sklearn 默认max_iter=100,但 MNIST 有 784 维特征和 10 个类别,优化器 100 轮根本不够,跑完你会发现警告刷屏,准确率也偏低。提到 300 是稳妥值,再往上提收益就很小了。multi_class='multinomial'表示用 Softmax 做多分类,比ovr(一对一)在这个场景下更合适,因为 MNIST 数字类别之间有相似性,多项逻辑斯蒂能捕捉类别间的概率竞争关系。

3.3 朴素贝叶斯:高斯分布假设下MNIST的奇特表现

bayes目录下的实现用的是GaussianNB,这个选择本身值得说一句:朴素贝叶斯假设特征之间独立,MNIST 的像素点显然不独立(相邻像素强相关),所以理论上它在该数据集上表现不会太好。但代码的价值恰恰在于给你一个“反直觉”的基线——跑出来你就会发现,它比想象中高,因为手写数字的像素分布有强先验。

from sklearn.naive_bayes import GaussianNB # var_smoothing 是平滑项,数值越大对噪声越容忍 gnb = GaussianNB(var_smoothing=1e-9) gnb.fit(x_train_sub, y_train_sub) val_acc = gnb.score(x_val, y_val) print(f'高斯朴素贝叶斯 验证集准确率: {val_acc:.4f}')

var_smoothing=1e-9是这里唯一的可调参数。它控制方差估计时的平滑量:如果某个像素在所有样本里取值几乎不变,方差会趋于 0,除零就会导致概率计算崩溃。调大平滑值能让模型更稳,但也会让边界变钝。我自己试过1e-2,准确率会掉零点几个点,所以1e-9是个不错的默认值。值得留意的是,高斯朴素贝叶斯在 MNIST 上的训练速度快到令人发指,几秒钟就完事,适合作为“先跑通整条流水线”的探路模型。

3.4 决策树:max_depth不设就是全量特征的黑匣子

decision_tree目录是最有意思的——里面除了main.py,还带了一个myTree.txt,是某次运行生成的树结构文本。打开这个文本你会看到一棵深不见底的树,一层层判断“第 234 个像素是否大于 0.5”。这恰恰是决策树在 MNIST 上的最大陷阱:不加深度限制,它会无限生长,把训练集背下来。

from sklearn.tree import DecisionTreeClassifier # max_depth 必须限制,否则树深随特征数膨胀,过拟合到没法看 dt = DecisionTreeClassifier( max_depth=12, # 经验值,8~16 之间都可以试 min_samples_leaf=4, # 叶子节点最少样本数,压制噪声分支 criterion='gini' ) dt.fit(x_train_sub, y_train_sub) train_acc = dt.score(x_train_sub, y_train_sub) val_acc = dt.score(x_val, y_val) print(f'决策树 训练集准确率: {train_acc:.4f}, 验证集准确率: {val_acc:.4f}')

max_depth=12是调参的核心:如果你用默认None,训练集准确率能到 99% 以上,但验证集会掉到 80% 上下,这就是典型的把噪声学进去了。min_samples_leaf=4的意思是每个叶子节点至少要有 4 个样本才允许生成,这个参数能过滤掉那些只对极个别样本有效的分裂。决策树在 MNIST 上你就别指望它拿最高分,它的作用有两个:一是作为集成学习的弱学习器底稿;二是让你直观看到“决策树如何逼近真实曲线”——每深一层,决策边界就多一次切分,但到后面全是在拟合噪声。

3.5 支持向量机:线性核与rbf核在784维上的取舍

svm目录里同样存在main.py和main1.py两个版本,主版本用的是SVC,并且同时演示了线性和 rbf 两种核。MNIST 原本就是 784 维,样本量 6 万,SVM 在这里是典型的“能做但别太贪”的模型——rbf 核准确率高,但训练时间以小时计。

from sklearn.svm import SVC # 先跑线性核,速度极快,拿一个基线分数 svm_linear = SVC(kernel='linear', C=1.0) svm_linear.fit(x_train_sub, y_train_sub) print(f'SVM 线性核 验证集准确率: {svm_linear.score(x_val, y_val):.4f}') # rbf 核准确率更高,但计算量爆炸,建议只在子集上验证 svm_rbf = SVC(kernel='rbf', gamma='scale', C=1.0) svm_rbf.fit(x_train_sub, y_train_sub) print(f'SVM rbf核 验证集准确率: {svm_rbf.score(x_val, y_val):.4f}')

gamma='scale'是 sklearn 自动根据特征数计算 gamma 值的方式,在 784 维下它会自动调小,避免高维空间里距离度量失衡。这段代码还藏着一个性能教训:rbf 核的 SVC 在 48000 个训练样本上 fit 一次可能要十几分钟,机器内存不够还会崩。我的建议是,在完整数据上跑线性核拿基线,rbf 核只在切出来的 5000 到 10000 个子集上验证效果,方向对再决定要不要全量跑。

3.6 感知机:原始形式与sklearn默认参数的不一致

preceptron目录下的实现很有诚意:main.py手写了一个感知机类,包括权重初始化和迭代更新,demo.py做可视化演示。感知机是最早的线性分类器,它和 SVM 的决策本质一样是找一条直线(超平面),差别在于感知机只在误分类点更新权重,没有间隔最大化的概念,所以最终解不唯一。

from sklearn.linear_model import Perceptron # Perceptron 在 sklearn 里本质上就是 SGD 分类器的一种特例 perc = Perceptron( max_iter=1000, # 感知机要迭代到收敛,100 轮不够 tol=1e-3, # 损失变化小于这个值就停 random_state=42 ) perc.fit(x_train_sub, y_train_sub) val_acc = perc.score(x_val, y_val) print(f'感知机 验证集准确率: {val_acc:.4f}')

max_iter=1000是手写版本和 sklearn 版本之间最容易出现差异的地方:手写代码里你控制的是“对整个训练集扫多少遍”,sklearn 的max_iter也是这个意思,但默认只有 100。MNIST 特征维度高,100 遍根本不足以让权重收敛到误差平稳区。tol=1e-3表示连续两次迭代的损失变化小于这个阈值就提前停止训练,它配合max_iter能省不少时间——实际跑的时候经常在几百轮就触发了提前停止。

3.7 AdaBoost:SAMME算法与弱学习器数量怎么组合

adaboost目录下的main1.py是集成学习的入口。AdaBoost 在 MNIST 上的经典用法是以决策树桩(深度为 1 的决策树)作弱学习器,通过反复调整样本权重,把多个弱分类器加权组合成一个强分类器。sklearn 的AdaBoostClassifier默认支持SAMME和SAMME.R两种算法,多分类场景下SAMME更通用。

from sklearn.ensemble import AdaBoostClassifier from sklearn.tree import DecisionTreeClassifier # 弱学习器用深度为 1 的决策树桩,这是 AdaBoost 的经典搭档 base_est = DecisionTreeClassifier(max_depth=1) ada = AdaBoostClassifier( estimator=base_est, n_estimators=200, # 弱学习器数量,越多拟合越强,但也容易过拟合 learning_rate=0.8, # 每轮的权重衰减,调低能增强泛化 algorithm='SAMME', random_state=42 ) ada.fit(x_train_sub, y_train_sub) val_acc = ada.score(x_val, y_val) print(f'AdaBoost 验证集准确率: {val_acc:.4f}')

n_estimators=200和learning_rate=0.8是一对配合参数:弱学习器多了,模型对训练集的拟合增强,但超过某个临界点后验证集准确率不再上升甚至下降;调低learning_rate会让每个弱学习器的权重更新幅度变小,需要更多树来补齐,但泛化能力通常更好。这份代码选 200 和 0.8 是个不错的起点,跑完看验证集曲线,如果准确率还在上升就加到 300 试试,如果已经持平就别再浪费训练时间了。algorithm='SAMME'在高版本 sklearn 里是默认值,写出来是为了提醒你:如果你用的是老版本,SAMME.R要求弱学习器能输出概率,深度为 1 的决策树桩概率输出不平滑,用SAMME更稳。

3.8 最大熵:从最大熵到逻辑斯蒂的等价关系

最大熵这个目录在别的 MNIST 教程里很少见,值得单独说。最大熵模型的思路是:在满足已知约束的条件下,选择熵最大的概率分布。当特征函数定义得当、用对数线性模型建模时,最大熵模型和逻辑斯蒂回归在数学上是等价的——这也是为什么max_shang目录里同时存在手写实现的max_Ent.py和调 sklearn 的main.py。

# max_Ent.py 的核心:用梯度下降迭代更新权重 def train_maxent(features, labels, lr=0.01, epochs=200): n_samples, n_features = features.shape n_classes = len(np.unique(labels)) # 权重初始化:全零或小随机数 weights = np.zeros((n_features, n_classes)) for epoch in range(epochs): # 线性部分 scores = features @ weights # softmax 得到概率分布 exp_scores = np.exp(scores - scores.max(axis=1, keepdims=True)) probs = exp_scores / exp_scores.sum(axis=1, keepdims=True) # 梯度:真实标签的 one-hot 减去预测概率 grad = features.T @ (probs - one_hot(labels, n_classes)) weights -= lr * grad / n_samples if epoch % 50 == 0: print(f'epoch {epoch} 完成') return weights

max_Ent.py里软最大函数scores.max(axis=1, keepdims=True)是数值稳定的关键:如果不先减去最大值,exp在大数值输入下可能溢出成inf。lr=0.01作为学习率在 200 轮内能收敛到一个不错的位置,但比 sklearn 的lbfgs要慢不少。我的建议是:手写版只用来理解最大熵的原理和梯度推导,要拿准确率还是跑main.py的LogisticRegression,两者结果一致正好反过来验证你的推导没写错。

4. 避坑:MNIST下载404、内存爆炸、收敛警告,逐个给你排查路径

4.1 torchvision下载MNIST报404或连接失败

现象:用torchvision.datasets.MNIST(root='./data', download=True)下载时报HTTP Error 404,或者连接超时卡死不动。

原因:MNIST 官方源在部分网络环境下访问不稳定,torchvision默认下载地址响应失败。这不是代码写错了,是网络链路问题。

解决:换成手动下载mnist.npz文件,放到项目的data目录,再用np.load读取,这一点正是这份源码包的做法。下载时注意不要解压,npz是压缩格式,代码里直接 load 就行。若np.load遇到编码问题,给np.load加上allow_pickle=False再读。

4.2 KNN把内存吃满,进程卡死

现象:knn.fit后进行knn.score时内存占用一直涨,最后进程被杀或卡死。

原因:KNN 不显式训练,但预测时要保存全部训练样本作为参考集。60000 个 784 维的数组本来就占内存,代码里如果用float64存储,每张图占 6272 字节,合计约 360 MB 的参考集,再加上 sklearn 内部距离矩阵的中间存储和n_jobs=-1的多进程复制,内存直接爆。

解决:数据读入后立刻astype('float32')把内存减半;algorithm优先选kd_tree或ball_tree,避免brute模式下生成全量距离矩阵;真遇到超大输入,用n_neighbors更小的值并配合chunk_size分批预测。

4.3 ConvergenceWarning 刷屏,模型准确率偏低

现象:跑逻辑斯蒂或最大熵时,控制台不停弹ConvergenceWarning: Maximum iterations reached,最终准确率也明显低于预期。

原因:sklearn 默认max_iter=100,在 784 维特征和 6 万级样本下,优化器 100 轮通常还没走到最优解附近。SVM 等其他迭代类模型同理。

解决:把max_iter提到 300 到 500。改完之后如果警告消失但准确率没明显变化,说明模型已经收敛,这时再加大max_iter没有意义。还有一个隐藏点:tol参数默认1e-4,如果你的数据没归一化,梯度量级大,tol更严格会触发更多迭代,归一化之后tol保持默认就好。

4.4 决策树训练集准确率99%,验证集却只有七成多

现象:决策树训练集准确率接近 1.0,验证集准确率明显下滑,两者差值超过 15 个百分点。

原因:决策树没限制深度,树把训练集里每个数字的笔画细节、噪声点全背下来了,这正好对应热搜里反复出现的“决策树如何逼近真实曲线”问题——它逼近过头了,逼近的不是真实曲线而是训练集的噪声。

解决:max_depth从 8 到 16 之间用 2 的步长网格搜索,同时设min_samples_leaf=4。这两个参数往大了调,训练集和验证集的准确率差距会缩小。注意一定要同时看两组数字,只盯验证集准确率也容易误判。

4.5 分类报告提示标签问题,准确率统计出错

现象:classification_report(y_test, y_pred)报ValueError: Classification metrics can't handle a mix of multilabel-indicator and multiclass targets,或者报告里类别信息乱掉。

原因:y_test是uint8或str类型,y_pred是int,两个数组 dtype 不一致,sklearn 的指标函数无法自动对齐。

解决:统一在数据装载阶段做y.astype(int),同时检查predict的结果是否经过.astype(int)。神经网络真实标签文件里这种情况非常常见,早期 dtype 不一致不报错,拖到分类报告阶段才暴露,很多人被这一下搞懵。

5. 把八种算法的结果横向拉齐:混淆矩阵、样本可视化与模型对比

5.1 用混淆矩阵看KNN和SVM的混淆集中在哪两类

跑完单个算法之后,下一步一定是横向对比。最简单有效的手段是每个模型生成一张混淆矩阵,看看预测错误集中在哪里——MNIST 上最经典的混淆是 4 和 9、3 和 8、7 和 2,这些数字形近,线性模型分不开很正常。

from sklearn.metrics import confusion_matrix, classification_report import seaborn as sns import matplotlib.pyplot as plt # y_pred 来自任意一个训练好的模型 cm = confusion_matrix(y_test, y_pred) plt.figure(figsize=(8, 6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues') plt.xlabel('预测标签') plt.ylabel('真实标签') plt.show() # 打印每个类别的精确率、召回率、F1 print(classification_report(y_test, y_pred))

fmt='d'让热力图显示整数计数而不是科学计数法,annot=True在每个格子里标出数值。分类报告里的macro avg是十个类别指标的平均,如果它明显低于accuracy,说明模型在某几个少数类别上特别差,光看总准确率会漏掉这个信息。

5.2 错误样本可视化:把预测错的图打印出来定位问题

混淆矩阵告诉你“哪里错了”,但不知道“错在什么图”上。把预测错误的样本打印成图片,是最直观的定位手段——你会发现很多错误样本连人眼都难分辨,模型能分对反而奇怪。

import math def show_misclassified(images, true_labels, pred_labels, num=10): errors = np.where(true_labels != pred_labels)[0] show_idx = errors[:num] # 取前 num 个错误样本 cols = 5 rows = math.ceil(num / cols) plt.figure(figsize=(cols * 2, rows * 2)) for i, idx in enumerate(show_idx): plt.subplot(rows, cols, i + 1) plt.imshow(images[idx].reshape(28, 28), cmap='gray') plt.title(f'真:{true_labels[idx]} 预:{pred_labels[idx]}') plt.axis('off') plt.tight_layout() plt.show() show_misclassified(x_test, y_test, y_pred, num=10)

errors = np.where(true_labels != pred_labels)[0]是所有预测错误的索引,errors[:num]取前 10 个。打印出来的图片如果看起来确实模糊、偏斜、有干扰线,就说明模型没问题;如果有几张图片人眼看很清楚但模型错了,那大概率是预处理或训练数据问题,比如训练时没做数据增强、图没对齐。

5.3 模型对比表:把八种算法的训练时间、准确率和适用场景拉齐

把八种算法跑完后,整理一张对比表是课程设计报告里必不可少的内容。我拆这份包时顺手跑了一轮,表格格式如下,你可以直接替换自己的运行结果:

算法训练时间(量级)内存占用验证集准确率(量级)适用场景
KNN秒级(训练)/ 分钟级(预测)高较高小数据集基线,容易快速出结果
逻辑斯蒂数十秒低较高需要训练快、可解释性强的场景
高斯朴素贝叶斯极快低中等先跑通流水线的探路模型
决策树秒级低偏低观察特征分裂过程、课程作业演示
线性 SVM秒级中较高高维稀疏数据的入门基线
RBF SVM小时级极高最高小样本高精度场景,不推荐全量跑
感知机数十秒低中等理解线性分类器的收敛过程
AdaBoost分钟级中中等偏高需要展示集成学习效果时

这张表的价值不只是给你一份答案,而是让你知道每种算法的定位:KNN 和逻辑斯蒂是“先拿基线”,决策树和感知机适合演示原理,SVM 和 AdaBoost 是“冲高分的选项但代价大”,朴素贝叶斯适合验证流水线是否通着。我在实际跑的时候会用脚本把每种算法的训练时间和准确率自动追加到 CSV 文件里,方便最后统一整理成报告表格。

6. 跑通之后把它变成工具:模型持久化与批量预测的一体化脚本

6.1 用joblib把最优模型和归一化参数一起存下来

训练完八种算法,挑出验证集准确率最高的那个模型,接下来的问题是:怎么把它保存下来,下次直接加载、不用重新训练。joblib是 sklearn 官方推荐的模型持久化方案,比pickle对 numpy 数组的压缩效果好得多。

import joblib # 保存模型的同时,把归一化的基准值也存进去 joblib.dump(best_model, 'models/best_model.pkl') joblib.dump({'scaler': 255.0}, 'models/norm_config.pkl') print('模型已保存到 models/best_model.pkl') # 下次加载 loaded_model = joblib.load('models/best_model.pkl') print(f'加载完成,模型类型: {type(loaded_model).__name__}')

保存时只存模型参数不存训练数据,joblib会自动处理 sklearn 模型内部复杂的对象结构。归一化基准单独存一份,这步很多人忽略——新图片进来自动预测时要除 255,但训练那次除 255 后模型已经学好了,如果你加载模型后忘了重新归一化,预测结果会全部乱掉。所以我把基准一起序列化,加载时一并取用。

6.2 新图片进来自动预测:一套处理28x28输入的完整函数

def predict_digit(image_array, model, norm_value=255.0): """输入任意形状的图片数组,自动转成 MNIST 标准格式并预测。""" if image_array.shape != (28, 28): from skimage.transform import resize image_array = resize(image_array, (28, 28), anti_aliasing=True) # 转灰度、拉平、归一化 img_flat = image_array.reshape(1, -1).astype('float32') / norm_value pred = model.predict(img_flat) return int(pred[0]) # 从项目目录里取一张测试图 sample = x_test[0] print(f'预测结果: {predict_digit(sample, loaded_model)}')

resize只在新图片不是 28 乘 28 时才执行,anti_aliasing=True可以避免缩放产生的锯齿伪影。函数内部统一做reshape(1, -1)表示“一条样本、全特征”,这样无论传入的是单张 28 乘 28 图还是已经拉平的一维数组,都能正确预测。

6.3 批量跑八种算法的自动化循环,给自己留一份可复现报告

models = { 'KNN': KNeighborsClassifier(n_neighbors=5, weights='distance', n_jobs=-1), 'Logistic': LogisticRegression(max_iter=300, solver='lbfgs', multi_class='multinomial'), 'BNB': GaussianNB(), 'Tree': DecisionTreeClassifier(max_depth=12, min_samples_leaf=4), 'LinearSVM': SVC(kernel='linear', C=1.0), 'Perceptron': Perceptron(max_iter=1000, tol=1e-3), 'AdaBoost': AdaBoostClassifier(n_estimators=200, learning_rate=0.8, algorithm='SAMME'), } results = [] for name, model in models.items(): model.fit(x_train_sub, y_train_sub) acc = model.score(x_val, y_val) results.append((name, acc)) print(f'{name}: {acc:.4f}') # 按准确率排序,输出前三名 results.sort(key=lambda x: x[1], reverse=True) print('\n前三名:') for rank, (name, acc) in enumerate(results[:3], start=1): print(f'{rank}. {name} - {acc:.4f}')

注意这个循环里没有把最大熵的代表放进去,因为手写版max_Ent.py的接口和 sklearn 的统一fit/predict风格不一致——这也是实际工程里的常态。把接口统一的七个模型放到循环里跑,手写模型单独跑,然后结果手动合并。批量循环跑一遍并打印前三名之后,把这轮结果存成日志文件,你手上的这包代码就成了一个可复现的实验报告模板。

拆完这份包、写完这份笔记,我自己的感想是:八种算法跑 MNIST 这件事本身不难,难的是第一次搭数据管道时各种怪坑——数据读不进来、内存爆掉、dtype 不一致、收敛警告刷屏。从那时起我每次处理新数据集,都会强制走一遍“先小样本跑通→确认 dtype 和归一化→再全量训练”的流程,这套习惯帮我省掉了无数次卡死在训练到一半的时间。这份源码包把八种算法的常见坑都踩过一遍了,希望帮到你。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/26 18:03:49

基于GAN的HDR图像合成与色调映射:从原理到工程实践

简介:面向图像处理与机器学习研究者的GAN实战资源包,聚焦高动态范围图像合成与色调映射全流程,适合具备一定深度学习基础、希望复现生成对抗网络在HDR领域应用的读者。压缩包共12个文件,包含6个Python脚本(负责数据加载…

作者头像 李华
网站建设 2026/9/26 18:03:20

电控岗秋招必备:10个可写进简历的开源项目详解

1. 先搞懂:招聘方在你的简历里翻什么投电控岗,简历上全是“学过电路、模电、数电、自动控制原理”,这基本等于没写。我当年也犯过这个错。校招HR一天筛几百份简历,电控方向的JD里永远写着“熟悉BLDC/PMSM电机控制”“熟悉PID等控制…

作者头像 李华
网站建设 2026/9/26 18:02:58

Java多线程与并发编程:从JMM到线程池的实战指南

Java多线程与并发,这个话题在面试和实战里被翻来覆去地问、反反复复地踩。很多人背了一堆八股文,从Thread到ThreadPoolExecutor,从synchronized到Lock,看似什么都懂,真到了线上排查问题、设计一个高并发接口的时候&…

作者头像 李华