尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

手写数字识别:用纯NumPy实现线性回归分类器

手写数字识别:用纯NumPy实现线性回归分类器 简介本资源是一套基于Python实现的手写数字识别系统的完整课程设计实践包面向计算机、人工智能或数字图像处理方向的初学者与高校学生解决从模型训练到实际图像识别的全流程实践问题。压缩包共16个文件含9幅28×28像素的黑白手写数字BMP样本图覆盖0–9、2个核心Python脚本训练与测试、2个CSV数据文件标签编码与模型权重、1份Word设计报告、1份README说明及LICENSE协议整体仅251KB轻量易部署。已有4324人学习下载适合作为机器学习入门项目快速上手。读者可直接运行代码复现多元线性回归模型的训练与推理过程结合设计报告理解多分类建模逻辑并利用提供的标准尺寸手绘图像样本验证识别效果具备完整的实验闭环与教学参考价值。1. 用28×28像素黑底白字图跑通一个真正能识别手写数字的线性回归模型你画一个歪歪扭扭的“7”保存成28×28 BMP格式、纯黑背景纯白笔迹丢进这个Python项目——它真能输出“7”而不是报错、卡死或返回随机数。这不是MNIST数据集上的玩具demo而是从零构建训练流程、固化权重、独立加载测试的完整闭环训练脚本生成myweight1.csv测试脚本直接读取该文件做前向推理全程不依赖TensorFlow/PyTorch只用NumPy和标准库。它专为课程设计场景打磨结构清晰训练/测试分离、可调试每层矩阵运算显式展开、可验证提供9张实拍手写图1.bmp~9.bmp及3.bmp且所有图像预处理逻辑全部内联在代码中——没有隐藏的PIL自动缩放、没有暗藏的归一化陷阱。适合刚学完线性代数与Python基础的学生动手复现也适合想快速验证分类器底层逻辑的工程师做最小可行性验证。2. 多元线性回归为何能胜任手写数字分类从数学本质到代码映射2.1 为什么选多元线性回归而非神经网络手写数字识别常被默认绑定CNN或全连接神经网络但本项目刻意回归最简模型——多元线性回归Multinomial Logistic Regression实际实现为带Softmax的线性分类器。其合理性在于MNIST类数据具有强线性可分性像素级特征虽高维但类别边界相对清晰且课程设计需暴露核心数学逻辑。相比深度模型线性回归的参数更新过程完全透明权重矩阵W形状为(784, 10)偏置向量b为(10,)输入x为(784,)向量输出z x·W b后经Softmax得10维概率分布。这种结构让每个像素对每个数字类别的贡献可直接追溯调试时能定位到“第327个像素权重异常”这类具体问题。而神经网络的隐层权重缺乏直观物理意义对初学者易形成黑箱认知。提示项目中train_label_hotencoding.csv即one-hot编码标签每行对应一张图的10维标签如数字3对应[0,0,0,1,0,0,0,0,0,0]这是线性分类器训练的必要输入格式不可省略或替换为整数标签。2.2 训练脚本手写数字的识别训练.py的三阶段拆解2.2.1 数据加载与预处理BMP→灰度→归一化→展平import numpy as np from PIL import Image def load_and_preprocess_image(filepath): # 1. 读取BMP并转为灰度避免彩色通道干扰 img Image.open(filepath).convert(L) # 2. 强制重采样为28x28关键原始图若非精确尺寸会破坏模型泛化 img img.resize((28, 28), Image.Resampling.LANCZOS) # 3. 转为numpy数组并归一化0-255 → 0-1线性回归对数值范围敏感 arr np.array(img) / 255.0 # 4. 反色处理BMP中数字为白色255背景为黑色0→ 模型习惯数字越亮特征越强 # 但本项目要求背景黑、数字白故无需反色若输入图是白底黑字则必须arr 1 - arr return arr.flatten() # 输出784维向量 # 示例加载训练集此处需自行构造train_images列表 train_images [load_and_preprocess_image(ftrain/{i}.bmp) for i in range(1000)] train_labels np.loadtxt(train_label_hotencoding.csv, delimiter,)逻辑说明resize使用LANCZOS抗锯齿算法比BILINEAR更保边缘锐度/255.0确保输入值域在[0,1]避免梯度爆炸flatten()将28×28矩阵压成784维向量与权重矩阵W的列数严格对齐。若跳过resize直接读取非28×28图后续矩阵乘法会因维度不匹配报错。2.2.2 损失函数与梯度推导交叉熵L2正则的闭式解项目未采用SGD迭代优化而是直接求解正规方程Normal Equation的近似解# X: (n_samples, 784) 输入矩阵, y: (n_samples, 10) one-hot标签 # 添加L2正则项λ0.01防止过拟合 lambda_reg 0.01 I np.eye(X.shape[1]) W np.linalg.inv(X.T X lambda_reg * I) X.T y np.savetxt(myweight1.csv, W, delimiter,)参数说明X.T X是784×784协方差矩阵lambda_reg * I为岭回归正则项np.linalg.inv求逆是计算瓶颈但对千量级样本仍可行X.T y本质是标签加权特征均值体现“每个数字类别的典型像素模式”。此解法比迭代更快且权重文件myweight1.csv可直接被测试脚本加载无需保存模型架构。2.2.3 权重文件myweight1.csv的结构验证生成的CSV文件含784行对应28×28像素、10列对应0~9数字类。可用以下代码验证其合理性W np.loadtxt(myweight1.csv, delimiter,) print(f权重矩阵形状: {W.shape}) # 应输出 (784, 10) print(f第0列数字0权重均值: {W[:,0].mean():.4f}) # 各列均值应接近0中心化 print(f第0列标准差: {W[:,0].std():.4f}) # 标准差反映特征重要性通常0.1若W[:,0].std()接近0说明数字0的判别特征未被学习需检查训练标签是否包含足够0样本或预处理是否错误。3. 测试脚本全流程执行从BMP输入到数字输出的端到端链路3.1 测试脚本手写数字的识别测试.py的四步执行逻辑3.1.1 加载固化权重与预处理单张图import numpy as np from PIL import Image # 1. 加载训练好的权重关键必须与训练时维度一致 W np.loadtxt(myweight1.csv, delimiter,) # 形状(784, 10) b np.zeros(10) # 本项目简化偏置为0实际可从训练中提取 # 2. 加载测试图以提供的1.bmp为例 test_img Image.open(1.bmp).convert(L).resize((28, 28), Image.Resampling.LANCZOS) x np.array(test_img) / 255.0 x x.flatten() # (784,) # 3. 前向传播z x·W b z x W b # 结果为(10,)向量 # 4. Softmax归一化为概率 exp_z np.exp(z - np.max(z)) # 减max防溢出 probs exp_z / np.sum(exp_z) # 5. 输出预测结果 pred_digit np.argmax(probs) confidence np.max(probs) print(f预测数字: {pred_digit}, 置信度: {confidence:.4f})逻辑说明x W是核心矩阵乘法符号在NumPy中明确表示矩阵乘非逐元素乘np.max(z)用于数值稳定避免exp(1000)导致infprobs各分量和为1可直接解释为概率。若confidence 0.6提示图像质量不足如笔迹过细、有噪点。3.1.2 批量测试多张图并生成结果表test_files [1.bmp, 2.bmp, 3.bmp, 4.bmp, 5.bmp, 6.bmp, 7.bmp, 8.bmp, 9.bmp] results [] for fname in test_files: img Image.open(fname).convert(L).resize((28, 28)) x np.array(img) / 255.0 z x.flatten() W probs np.exp(z - np.max(z)) / np.sum(np.exp(z - np.max(z))) pred np.argmax(probs) results.append([fname, pred, f{np.max(probs):.4f}]) # 输出Markdown表格便于对比 print(| 图像 | 预测数字 | 置信度 |) print(|------|----------|--------|) for r in results: print(f| {r[0]} | {r[1]} | {r[2]} |)运行后可得如下结果示例图像预测数字置信度1.bmp10.92152.bmp20.87333.bmp30.8921.........若某张图如8.bmp预测为3需检查该图是否书写变形如8的上半圆闭合不全被误判为3。3.2 关键参数调试表影响识别准确率的三大变量参数可调范围推荐值效果说明验证方法lambda_regL2正则系数0.001 ~ 0.10.01值过大导致权重趋近0所有预测概率均等过小则过拟合训练集修改后重新训练观察测试集置信度方差理想值下各数字置信度0.8且方差0.05图像二值化阈值0.1 ~ 0.90.5当手写图对比度低时降低阈值如0.3可增强笔迹在预处理中添加arr (arr 0.3).astype(float)再测试1.bmpSoftmax温度系数T0.5 ~ 2.01.0T1使概率分布更尖锐高置信度T1更平滑降低过拟合风险修改probs exp_z / np.sum(exp_z)为probs exp_z**T / np.sum(exp_z**T)注意所有参数调整必须同步修改训练与测试脚本否则权重与推理逻辑不匹配。4. 排查常见失败场景从报错信息定位根本原因4.1 维度不匹配错误的三层诊断法当运行测试脚本报错ValueError: shapes (784,) and (10,784) not aligned说明矩阵乘法维度错误。按顺序排查检查权重文件维度head -n 5 myweight1.csv | csvlook # 若无csvlook用Excel打开看行列数正确应为784行×10列。若为10行×784列是训练时W存储方向错误应存为W.T需在训练脚本中改为np.savetxt(myweight1.csv, W.T, delimiter,)。验证输入向量长度x np.array(Image.open(1.bmp).resize((28,28))).flatten() print(len(x)) # 必须输出784若输出783或785说明resize未生效需确认PIL版本旧版PIL可能忽略resample参数。确认矩阵乘法顺序错误写法z W xW为784×10x为784×1 → 不可乘正确写法z x Wx为1×784W为784×10 → 输出1×104.2 低置信度问题的图像质量根因分析若所有测试图置信度均0.5大概率是图像预处理缺陷。用以下代码可视化像素分布import matplotlib.pyplot as plt # 加载1.bmp并统计像素值分布 img np.array(Image.open(1.bmp).convert(L)) plt.hist(img.ravel(), bins256, range(0,255), alpha0.7) plt.xlabel(Pixel Value) plt.ylabel(Frequency) plt.title(Histogram of 1.bmp) plt.show()典型问题与修复直方图峰值在0附近但存在大量中间灰度值50~200→ 图像未充分二值化添加arr (arr 128).astype(float)直方图双峰0和255为主但有宽峰在100左右→ 手写图有阴影或扫描噪声需中值滤波from scipy.ndimage import median_filter; arr median_filter(arr, size3)直方图仅集中在0-10区间→ 图像过暗需全局增亮arr np.clip(arr * 1.5, 0, 255)。4.3 Windows画图软件绘制规范清单为确保输入图符合要求必须遵守新建画布尺寸设为28×28像素画图→文件→属性→自定义大小使用铅笔工具非刷子避免羽化颜色选择纯白RGB 255,255,255背景保持纯黑RGB 0,0,0数字居中绘制笔画宽度≥2像素过细则像素丢失保存为单色BMP画图→文件→另存为→BMP图片→在“另存为”对话框底部选择“单色位图”。提示若用Windows 11新版画图需先“另存为PNG”再用PIL转换Image.open(input.png).convert(1).save(1.bmp)因新版画图不支持直接存单色BMP。5. 进阶技巧用热力图可视化模型“看到”的数字特征5.1 构造数字类别的显著性热力图线性回归权重矩阵W的每一列代表一个数字类别的“重要像素”分布。以数字5为例提取其权重并重塑为28×28热力图import matplotlib.pyplot as plt # 加载权重 W np.loadtxt(myweight1.csv, delimiter,) # 取数字5的权重索引4因0~9对应列0~9 weights_5 W[:, 4].reshape(28, 28) # 绘制热力图 plt.figure(figsize(6,6)) plt.imshow(weights_5, cmapRdBu_r, vmin-0.1, vmax0.1) plt.colorbar(labelWeight Value) plt.title(Digit 5: Pixel Importance (RedPositive, BlueNegative)) plt.axis(off) plt.savefig(digit5_heatmap.png, bbox_inchestight) plt.show()参数说明cmapRdBu_r使正权重红色表示该像素亮起时倾向预测为5负权重蓝色表示该像素亮起时抑制预测为5vmin/vmax限制色阶范围避免单个异常值主导颜色映射。5.2 解读热力图指导手写优化观察生成的digit5_heatmap.png典型模式为顶部横线区域第3~5行权重为正→ 模型认为5的上横线是关键特征左上角第2行第5列权重为负→ 此处若出现白点如书写时连笔到左上会降低5的概率右下角弧线第20~25行第15~22列权重显著为正→ 5的下半圆是强判据。据此可指导用户写5时务必强化上横线与右下弧线避免左上角沾墨。同理分析digit8_heatmap.png会发现两个同心圆环区域权重最高若手写8的上下环不闭合对应环区权重将衰减导致置信度下降。5.3 将热力图集成到测试流程中在测试脚本末尾添加自动热力图生成功能def generate_heatmap_for_prediction(image_path, digit_class): W np.loadtxt(myweight1.csv, delimiter,) weights W[:, digit_class].reshape(28, 28) plt.figure(figsize(4,4)) plt.imshow(weights, cmapRdBu_r, vmin-0.05, vmax0.05) plt.title(fFeature Map for Predicted Digit {digit_class}) plt.axis(off) plt.savefig(fheatmap_{digit_class}_{image_path.split(.)[0]}.png) plt.close() # 在预测后调用 pred_digit np.argmax(probs) generate_heatmap_for_prediction(1.bmp, pred_digit)运行后生成heatmap_1_1.png直观展示模型为何认定这张图是“1”——通常显示垂直中轴线权重最高两侧为负权重印证了“1”的典型结构。本文还有配套的精品资源点击获取
返回列表