9.5 机器学习案例:手写数字分类


9.5 机器学习案例:手写数字分类

从零训练一个识别 0–9 的全连接网络:造数据集、划训练测试集、组网训练、评估准确率、存检查点。综合第 6 章 Flux 与第 4 章的 JLD2 检查点方案。

数据准备

没有现成数据集也能做:用带噪声的模板生成"伪数字"图像(真实项目换 MLDatasets 的 MNIST 即可):

] add Flux JLD2 using Flux, Statistics, Random Random.seed!(0) function make_sample(digit) # 5x5 像素图:在基础图案上加噪声 base = zeros(5, 5) base[2:4, 2:4] .= 1.0 base[digit % 3 + 2, 1] = 1.0 clamp.(base .+ randn(5, 5) .* 0.2, 0, 1) end X = hcat([vec(make_sample(d)) for d in rand(0:9, 1000)]...) y = Flux.onehotbatch(rand(0:9, 1000), 0:9) # 标签 one-hot # 留出法划分:80% 训练 / 20% 测试 idx = 1:800 Xtr, ytr = X[:, idx], y[:, idx] Xte, yte = X[:, 800 .+ (1:200)], y[:, 800 .+ (1:200)]

组网与训练

model = Chain( Dense(25, 64, relu), Dense(64, 10), ) loss(x, yy) = Flux.logitcrossentropy(model(x), yy) opt = Adam(0.005) for epoch in 1:100 Flux.train!(loss, params(model), [(Xtr, ytr)], opt) if epoch % 20 == 0 println("轮 $epoch 训练损失 = ", round(loss(Xtr, ytr), digits = 4)) end end

评估与错误分析

pred = Flux.onecold(model(Xte)) .- 1 truth = Flux.onecold(yte) .- 1 acc = mean(pred .== truth) println("测试集准确率:$(round(acc * 100, digits=1))%") # 错误分析:看模型把谁认成了谁 using StatsBase errors = countmap(zip(pred[ pred .!= truth ], truth[ pred .!= truth ]))

准确率只是一个数字,errors 里"3 被认成 8"这类模式才是改进方向的线索。

检查点:训练不怕中断

using JLD2 @save "model_ckpt.jld2" model epoch=100 acc=acc @load "model_ckpt.jld2" model

ML 项目五阶段与产出物

阶段 动作 产出物
数据 造/读数据、划分 Xtr/Xte/ytr/yte
建模 Chain 组网 model
训练 train! 循环看损失 收敛的权重
评估 准确率 + 错误分析 acc 与混淆模式
存档 JLD2 检查点 可复活的模型文件

⚠️ 常见坑:用训练集上的准确率报喜——它会随训练轮数单调上升而毫无意义。报告里只认测试集数字;进一步严谨做法是交叉验证。

💡 关键直觉:先让最小网络在一个极小数据集上"过拟合成功"(训练准确率接近 100%),确认管道无 bug,再换真数据与更大网络。管道本身有错时,任何调参都是玄学。

结果解读与改进迭代

拿到一个准确率数字后,改进方向从错误分析里找。假设错误集中在两三类数字互相混淆,按顺序尝试四个手段并逐个记录测试集变化:其一,加宽网络(Dense(25, 128, relu))——容量不足时最直接;其二,训练轮数加倍并观察损失是否仍在降——没降完就是欠训练;其三,给输入加更强的数据增广(随机遮挡一两格像素)——提升对噪声的鲁棒性;其四,换学习率调度(前期 0.01 后期 0.001)——收敛后期大步长会来回震荡。每个手段单独改、单独测,改动记录成一张表:

# 实验记录骨架:每次改动追加一行,结论有据可查 experiments = [ (改动 = "基线", acc = acc), (改动 = "宽度 64→128", acc = acc2), (改动 = "轮数 100→300", acc = acc3), ]

这条"单变量改动 + 测试集数字 + 记录"的纪律,比任何调参玄学都可靠;多数时候你会发现真正有效的改动只有一两个,其余全是噪音波动——不做对照实验根本分不出来。变式:把五阶段管道与 9.2 的函数化思路合流,训练入口接受超参数字典、输出实验记录,一个可复现的实验框架就成型了,这正是科研代码与工程代码的分水岭。

复现与随机性管理

机器学习案例要成为"可复现的实验",随机性必须被驯服。三个动作:固定所有随机源——数据生成、参数初始化、批内采样各自设独立的随机种子(Random.seed! 分阶段调用),任何一次重跑才能逐位一致;记录环境——Flux 版本、Julia 版本、超参数一起写进实验记录,训练结果才有归属;隔离实验——每个改动一个独立分支或目录,避免"我也不知道哪个改动起了作用"的窘境。这套纪律与 8.3 的测试文档一脉相承,区别在于机器学习的"正确性"没有确定答案,只能靠"可复现的对比"来逼近真相。做完这些,本案例的检查点文件里存的就不仅是模型权重,还有 seedacc 与超参数的完整快照——一个真正意义上的实验存档。

本节要点回顾

  • 五阶段管道:数据、建模、训练、评估、存档,每阶段有明确产出物;
  • one-hot 进、onecold 出,中间是交叉熵;
  • 测试集数字才算数,错误分析比总准确率更有指导性;
  • JLD2 检查点让长训练可中断可续跑,工程标配;
  • 单变量改动 + 对照记录,调参从玄学变成实验科学。

案例用伪数据是刻意为之的安排:它把"管道是否正确"与"数据是否够好"两个问题解耦了。管道在简单数据上跑通并过拟合成功后,换真实数据集只需要替换数据准备段,其余四阶段代码原封不动——这种"先简易数据验管道、再真数据出结果"的顺序,能帮你把 bug 定位范围从整个项目缩小到数据本身。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U