最近我完成了一个小项目:MNIST Neural Network Visualizer。
它是一个运行在浏览器里的手写数字识别器:在黑色画布上写下一个数字,模型会随着笔迹实时更新判断结果。与此同时,页面还会把输入层、隐藏层神经元、层间连接以及 0~9 的概率分布一起画出来。
我希望它不只是告诉你“这是数字 6”,还能够把神经网络得到这个答案的过程直观地展示出来。
这个项目做了什么
项目使用的是一个结构很简单的多层感知机:
784 → 32 → 16 → 10- 784 个输入,对应一张
28 × 28的灰度图 - 两个隐藏层,分别包含 32 和 16 个神经元
- 10 个输出,对应数字 0~9 的概率
- 隐藏层使用 ReLU,输出层使用 softmax
- 总计 25,818 个参数,MNIST 测试集准确率为 96.9%
整个训练过程只依赖 NumPy,包括 He 初始化、前向传播、小批量随机梯度下降和反向传播,没有使用 PyTorch、TensorFlow 等深度学习框架。
训练完成后,train.py 会将权重和偏置导出到 model.json。浏览器通过原生 JavaScript 读取模型并执行前向推理,因此网页运行时不需要后端,也不会把用户写下的数字上传到服务器。
展示网络内部
普通的手写数字识别 Demo 往往只给出一个预测结果。这个项目则把一次推理拆成了几个同步变化的区域:
- 书写区:支持鼠标和触屏输入
- 输入层:以
28 × 28热力图显示模型真正接收到的像素 - 隐藏层:显示各个神经元当前的激活程度
- 权重连线:绿色代表正权重,红色代表负权重,线条粗细随激活强度变化
- 输出层和概率条:实时显示模型对 0~9 的置信度,并高亮最终预测
书写过程中,推理与绘制通过 requestAnimationFrame 调度。每增加一笔,输入都会重新预处理并完成一次前向传播,所以能看到模型的判断如何从不确定逐渐收敛。
最关键的部分:让手写输入接近 MNIST
模型准确率并不只取决于网络本身。网页画布上的笔迹如果与 MNIST 数据的分布差异太大,即使模型在测试集上表现很好,实际书写时也可能频繁识别错误。
因此,app.js 中的 preprocess() 会依次完成:
- 找到笔迹的包围盒
- 将笔迹放大并居中到
112 × 112的超采样画布 - 通过形态学处理统一笔画粗细,使最终笔画接近 MNIST 的约 2.4 像素
- 使用
4 × 4均值采样缩小回28 × 28 - 根据灰度质心再次平移居中
其中,笔画粗细归一化尤其重要。直接把一幅很大的手写图压缩到 28 × 28,笔画可能变得过细甚至断裂。例如数字 6 底部的圆环一旦断开,就很容易被模型误认为 5。先在高分辨率画布中统一粗细,再降采样,可以明显改善这种情况。
从零训练的过程
训练脚本的核心逻辑很短,却包含了一个神经网络训练所需的完整步骤。
前向传播依次计算:
输入 → 线性变换 → ReLU → 线性变换 → ReLU → 线性变换 → softmax训练时,softmax 输出与 one-hot 标签的误差从输出层向前传递,逐层求出权重和偏置的梯度,再用学习率进行更新。默认配置为:
net = MLP([784, 32, 16, 10])net.train(xtr, ytr, xte, yte, epochs=25, bs=128, lr=0.15)网络规模是有意控制得比较小的。一方面,它在普通 CPU 上几十秒即可完成 25 个 epoch;另一方面,32 和 16 个隐藏单元能够全部清楚地画在页面上。这里追求的不是极限准确率,而是在识别能力、运行速度和可解释的视觉效果之间取得平衡。
导出时,权重会转置为 [out, in] 的形式,与前端逐个输出神经元计算点积的方式对应。这样 Python 训练端与 JavaScript 推理端使用的是完全相同的一组参数。
在本地运行
仓库已经包含训练好的 model.json,不需要先训练模型。克隆项目并启动一个静态文件服务器即可:
git clone https://github.com/Bihrys/mnist-neural-net-visualizer.gitcd mnist-neural-net-visualizerpython -m http.server 8000然后访问 http://localhost:8000。
需要注意,不能直接双击 index.html 打开。页面通过 fetch() 加载 model.json,浏览器会阻止网页在 file:// 协议下读取该文件,因此必须通过 HTTP 访问。
如果想亲自重新训练,只需安装 NumPy 后运行:
pip install numpypython train.py脚本会读取项目目录中的 mnist.npz,训练结束后重新生成 model.json。
项目结构
mnist-neural-net-visualizer/├── index.html # 页面结构├── style.css # 页面样式├── app.js # 画布、预处理、推理与可视化├── train.py # 纯 NumPy 训练脚本├── model.json # 已训练的模型参数└── mnist.npz # MNIST 训练与测试数据这个项目让我真正把“训练一个模型”和“让模型成为可交互的产品”连接了起来:从 NumPy 中的矩阵运算,到 JSON 模型格式,再到浏览器里的实时推理和可视化,每一步都可以直接打开源码查看。
如果你也对神经网络的工作过程感兴趣,欢迎在 GitHub 上体验、阅读源码或提出建议。