news 2026/8/4 11:53:09

Python实现简单神经网络:从零开始手写数字识别

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Python实现简单神经网络:从零开始手写数字识别

1. 项目概述

"根据Python基础语法完成一个简单的深度学习模拟"这个项目非常适合刚掌握Python基础语法,想要迈入深度学习领域的初学者。通过这个实践,你不仅能巩固Python编程基础,还能直观理解深度学习中最核心的概念——前向传播和反向传播。

我在指导新人学习时发现,很多初学者在刚接触深度学习时容易被复杂的数学公式吓退。实际上,用不到100行Python代码,就能实现一个完整的神经网络训练过程。这个项目将使用纯Python(不依赖任何深度学习框架)来实现一个能识别手写数字的简单神经网络。

2. 核心概念解析

2.1 神经网络基本原理

神经网络本质上是一个由多个函数组成的复合函数。以识别手写数字为例,输入是28×28=784个像素值,输出是0-9十个数字的概率分布。中间的"隐藏层"就是我们要通过训练来确定的参数。

这里我们实现最简单的全连接网络:

  • 输入层:784个神经元(对应28×28图像)
  • 隐藏层:16个神经元(这个数量可以调整)
  • 输出层:10个神经元(对应0-9十个数字)

2.2 关键数学概念

  1. Sigmoid激活函数:将神经元的输出压缩到(0,1)区间

    def sigmoid(x): return 1 / (1 + np.exp(-x))
  2. Softmax函数:将输出层的值转换为概率分布

    def softmax(x): exps = np.exp(x - np.max(x)) return exps / np.sum(exps)
  3. 交叉熵损失函数:衡量预测结果与真实标签的差距

    def cross_entropy(y_pred, y_true): return -np.sum(y_true * np.log(y_pred))

3. 完整实现步骤

3.1 数据准备

我们使用MNIST数据集,包含60000张手写数字图片。可以通过以下方式加载:

from tensorflow.keras.datasets import mnist (train_images, train_labels), (test_images, test_labels) = mnist.load_data() # 数据预处理 train_images = train_images.reshape((60000, 28*28)) / 255.0 test_images = test_images.reshape((10000, 28*28)) / 255.0 # 将标签转为one-hot编码 train_labels = np.eye(10)[train_labels] test_labels = np.eye(10)[test_labels]

3.2 网络初始化

class NeuralNetwork: def __init__(self, input_size, hidden_size, output_size): # 初始化权重和偏置 self.W1 = np.random.randn(input_size, hidden_size) * 0.01 self.b1 = np.zeros(hidden_size) self.W2 = np.random.randn(hidden_size, output_size) * 0.01 self.b2 = np.zeros(output_size)

3.3 前向传播实现

def forward(self, X): # 第一层计算 self.z1 = np.dot(X, self.W1) + self.b1 self.a1 = sigmoid(self.z1) # 输出层计算 self.z2 = np.dot(self.a1, self.W2) + self.b2 self.a2 = softmax(self.z2) return self.a2

3.4 反向传播实现

这是最核心的部分,计算损失函数对各个参数的梯度:

def backward(self, X, y, output): # 输出层误差 delta2 = output - y # 隐藏层误差 delta1 = np.dot(delta2, self.W2.T) * (self.a1 * (1 - self.a1)) # 计算梯度 dW2 = np.dot(self.a1.T, delta2) db2 = np.sum(delta2, axis=0) dW1 = np.dot(X.T, delta1) db1 = np.sum(delta1, axis=0) return dW1, db1, dW2, db2

3.5 参数更新

使用梯度下降法更新参数:

def update_params(self, dW1, db1, dW2, db2, lr=0.01): self.W1 -= lr * dW1 self.b1 -= lr * db1 self.W2 -= lr * dW2 self.b2 -= lr * db2

4. 训练过程与评估

4.1 训练循环实现

def train(self, X, y, epochs=10, lr=0.01): for epoch in range(epochs): # 前向传播 output = self.forward(X) # 计算损失 loss = cross_entropy(output, y) # 反向传播 dW1, db1, dW2, db2 = self.backward(X, y, output) # 更新参数 self.update_params(dW1, db1, dW2, db2, lr) if epoch % 10 == 0: print(f"Epoch {epoch}, Loss: {loss:.4f}")

4.2 模型评估

def evaluate(self, X_test, y_test): predictions = self.forward(X_test) predicted_labels = np.argmax(predictions, axis=1) true_labels = np.argmax(y_test, axis=1) accuracy = np.mean(predicted_labels == true_labels) print(f"Test Accuracy: {accuracy*100:.2f}%")

5. 完整代码示例

import numpy as np from tensorflow.keras.datasets import mnist # 激活函数和损失函数 def sigmoid(x): return 1 / (1 + np.exp(-x)) def softmax(x): exps = np.exp(x - np.max(x)) return exps / np.sum(exps, axis=1, keepdims=True) def cross_entropy(y_pred, y_true): return -np.sum(y_true * np.log(y_pred + 1e-8)) class NeuralNetwork: def __init__(self, input_size, hidden_size, output_size): self.W1 = np.random.randn(input_size, hidden_size) * 0.01 self.b1 = np.zeros(hidden_size) self.W2 = np.random.randn(hidden_size, output_size) * 0.01 self.b2 = np.zeros(output_size) def forward(self, X): self.z1 = np.dot(X, self.W1) + self.b1 self.a1 = sigmoid(self.z1) self.z2 = np.dot(self.a1, self.W2) + self.b2 self.a2 = softmax(self.z2) return self.a2 def backward(self, X, y, output): delta2 = output - y delta1 = np.dot(delta2, self.W2.T) * (self.a1 * (1 - self.a1)) dW2 = np.dot(self.a1.T, delta2) db2 = np.sum(delta2, axis=0) dW1 = np.dot(X.T, delta1) db1 = np.sum(delta1, axis=0) return dW1, db1, dW2, db2 def update_params(self, dW1, db1, dW2, db2, lr=0.01): self.W1 -= lr * dW1 self.b1 -= lr * db1 self.W2 -= lr * dW2 self.b2 -= lr * db2 def train(self, X, y, epochs=10, lr=0.01): for epoch in range(epochs): output = self.forward(X) loss = cross_entropy(output, y) dW1, db1, dW2, db2 = self.backward(X, y, output) self.update_params(dW1, db1, dW2, db2, lr) if epoch % 10 == 0: print(f"Epoch {epoch}, Loss: {loss:.4f}") def evaluate(self, X_test, y_test): predictions = self.forward(X_test) predicted_labels = np.argmax(predictions, axis=1) true_labels = np.argmax(y_test, axis=1) accuracy = np.mean(predicted_labels == true_labels) print(f"Test Accuracy: {accuracy*100:.2f}%") # 加载数据 (train_images, train_labels), (test_images, test_labels) = mnist.load_data() train_images = train_images.reshape((60000, 28*28)) / 255.0 test_images = test_images.reshape((10000, 28*28)) / 255.0 train_labels = np.eye(10)[train_labels] test_labels = np.eye(10)[test_labels] # 创建并训练模型 nn = NeuralNetwork(784, 16, 10) nn.train(train_images[:1000], train_labels[:1000], epochs=100, lr=0.1) nn.evaluate(test_images[:1000], test_labels[:1000])

6. 常见问题与优化建议

6.1 梯度消失问题

当网络层数增加时,Sigmoid函数容易导致梯度消失。这是因为Sigmoid的导数最大值为0.25,在反向传播时梯度会逐层衰减。解决方案:

  • 使用ReLU激活函数替代Sigmoid
  • 使用更先进的初始化方法(如He初始化)

6.2 学习率选择

学习率太大可能导致震荡不收敛,太小则训练缓慢。建议:

  • 初始尝试0.01-0.1范围
  • 实现学习率衰减策略
  • 考虑使用自适应优化器(如Adam)

6.3 过拟合处理

当训练集准确率远高于测试集时,说明模型过拟合。解决方法:

  • 增加训练数据量
  • 添加L2正则化
  • 实现Dropout技术

6.4 性能优化技巧

  1. 向量化计算:避免使用for循环,充分利用NumPy的矩阵运算
  2. 批量训练:不要一次性加载所有数据,实现mini-batch训练
  3. GPU加速:对于更大规模的数据,考虑使用CUDA加速

7. 项目扩展方向

这个基础项目可以进一步扩展:

  1. 增加隐藏层数量,实现更深网络
  2. 实现卷积神经网络(CNN)处理图像
  3. 添加批归一化(BatchNorm)层
  4. 实现自动求导功能
  5. 构建更复杂的损失函数

我在实际教学中发现,通过这个基础实现,学生能更深刻地理解PyTorch/TensorFlow等框架背后的工作原理。当你知道每一行代码在做什么时,使用高级框架就会更加得心应手。

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

企业开票信息API接口开发与优化实践

1. 企业开票信息查询API接口概述 发票抬头查询API接口是企业财税数字化进程中不可或缺的基础服务组件。作为连接企业信息系统与税务数据的桥梁,这类接口解决了传统开票流程中信息核验效率低下的痛点。我在为多家企业实施财税系统集成时发现,超过60%的票据…

作者头像 李华
网站建设 2026/8/4 11:52:07

DMA技术详解:从原理到STM32串口DMA实战应用

1. 先别被缩写吓到,DMA解决的是CPU“打杂”问题 DMA,全称Direct Memory Access,直接存储器访问。这个名字听起来很技术,但它的核心目标非常直接: 让CPU从繁重的数据搬运工作中解放出来 。 想象一个场景:…

作者头像 李华
网站建设 2026/8/4 11:51:23

C++函数封装进阶:从参数设计到模板实战的工程实践

1. 从“能用”到“好用”:为什么函数封装是C进阶的必经之路如果你写过一些C代码,可能已经习惯了把功能一股脑塞进main函数,或者随意定义几个函数来处理特定任务。这当然能跑起来,但代码很快就会变得像一团乱麻——修改一个地方&am…

作者头像 李华
网站建设 2026/8/4 11:50:38

Godot引擎RPG开发全流程:从架构设计到核心系统实现

1. 项目概述:为什么选择Godot来构建你的RPG? 如果你正在寻找一个既能让你天马行空地构思世界观,又不会在技术实现上把你卡死的游戏引擎来制作一款RPG,那么Godot引擎绝对是一个值得你投入时间研究的选项。我最初接触Godot&#xff…

作者头像 李华
网站建设 2026/8/4 11:49:40

Linux中文字体安装指南:从乱码到完美显示

1. 为什么在Linux上安装中文字体库是个“刚需”? 如果你在Linux上打开一个中文文档、浏览一个中文网页,或者运行一个带中文界面的软件,看到的却是一堆“口口口”或者乱码方块,那感觉就像对着满汉全席却只能闻味儿。这十有八九是系…

作者头像 李华
网站建设 2026/8/4 11:49:10

Ubuntu 22.04安装NS3网络模拟器:从依赖配置到编译运行的完整指南

1. 为什么在Ubuntu上安装NS3依然是个“技术活”?如果你正在学习网络协议、准备进行网络仿真研究,或者想复现一篇顶会论文里的实验,NS3(Network Simulator 3)大概率是你绕不开的一个工具。作为一个开源的、离散事件驱动…

作者头像 李华