基于 SVM(支持向量机)的手写数字识别

📌 主要步骤

  1. 安装必要的库
  2. 加载数据集(MNIST 手写数字)
  3. 数据预处理
  4. 划分训练集和测试集
  5. 训练 SVM 模型
  6. 评估模型
  7. 预测并可视化结果

1️⃣ 安装必要的库

在开始之前,请确保你的环境安装了以下库:

python 复制代码
pip install numpy pandas scikit-learn matplotlib

2️⃣ 加载数据集

我们将使用 scikit-learn 自带的 digits 数据集,它包含 0-9 的手写数字,每张图片是 8x8 像素的灰度图

python 复制代码
from sklearn import datasets
import matplotlib.pyplot as plt

# 加载手写数字数据集
digits = datasets.load_digits()

# 显示数据集信息
print("数据集形状:", digits.data.shape)
print("标签类别:", digits.target_names)

# 显示前 5 条数据
print("前 5 个标签:", digits.target[:5])

# 可视化前 10 张手写数字
fig, axes = plt.subplots(1, 10, figsize=(10, 3))
for i, ax in enumerate(axes):
    ax.imshow(digits.images[i], cmap='gray')
    ax.set_title(f"Label: {digits.target[i]}")
    ax.axis("off")
plt.show()

3️⃣ 数据预处理

将图片数据转换为一维数组 (从 8x8=64 变成 64 维 的特征向量),以便进行训练。

python 复制代码
import pandas as pd

# 转换为 DataFrame 以便查看
df = pd.DataFrame(digits.data)
df['label'] = digits.target

# 显示前 5 行数据
print(df.head())

4️⃣ 划分训练集和测试集

将数据集划分为 80% 训练集20% 测试集

python 复制代码
from sklearn.model_selection import train_test_split

# 划分数据
X_train, X_test, y_train, y_test = train_test_split(digits.data, digits.target, test_size=0.2, random_state=42)

# 打印数据集大小
print(f"训练集样本数: {len(X_train)}, 测试集样本数: {len(X_test)}")

5️⃣ 训练 SVM(支持向量机)模型

SVM(支持向量机)是一个非常强大的分类算法,在手写数字识别任务中表现优秀。

python 复制代码
from sklearn.svm import SVC

# 创建 SVM 模型
svm_model = SVC(kernel='linear')  # 使用线性核函数

# 训练模型
svm_model.fit(X_train, y_train)

print("SVM 模型训练完成!")

6️⃣ 评估模型

计算模型在测试集上的准确率。

python 复制代码
from sklearn.metrics import accuracy_score

# 预测测试集
y_pred = svm_model.predict(X_test)

# 计算准确率
accuracy = accuracy_score(y_test, y_pred)
print(f"模型准确率: {accuracy:.2f}")

7️⃣ 预测并可视化结果

我们从测试集中选择 10 张手写数字,进行预测并可视化。

python 复制代码
import numpy as np

# 选择前 10 个测试样本
sample_images = X_test[:10]
sample_labels = y_test[:10]

# 进行预测
predictions = svm_model.predict(sample_images)

# 可视化预测结果
fig, axes = plt.subplots(1, 10, figsize=(10, 3))
for i, ax in enumerate(axes):
    ax.imshow(sample_images[i].reshape(8, 8), cmap='gray')
    ax.set_title(f"P:{predictions[i]}\nT:{sample_labels[i]}")
    ax.axis("off")

plt.show()

项目总结

通过这个项目,我们完成了一个 机器学习分类任务 : ✅ 加载 MNIST 数据集

✅ 数据预处理(转换 8x8 图片到 64 维特征向量)

✅ 划分数据集

✅ 训练 SVM 分类器

✅ 评估分类准确率

✅ 预测并可视化结果

相关推荐
南京兴帝文化传媒有限公司1 分钟前
本地生活服务商户GEO优化技术实践:AI大模型收录机制与地图POI权重算法拆解
大数据·人工智能·算法·生活·geo优化实操·csdn运营技巧·ai内容收录
2601_962293538 分钟前
人工智能 & 神经网络完整入门路线(零基础可走,分阶段)
人工智能·python·深度学习·神经网络·机器学习
拓海SEO外贸28 分钟前
关键词聚类(Keyword Clustering)怎么做
人工智能·机器学习·聚类
纪伊路上盛名在29 分钟前
MISATO:基于结构的药物发现的蛋白-配体复合物机器学习数据集
深度学习·机器学习·数据集·分子动力学模拟·蛋白质结构·药物发现·蛋白质-配体互作
linx2951 小时前
单元九 · 零基础路线图-第 7–9 周·类与 RAII
c语言·开发语言·数据结构·c++·嵌入式硬件·算法
布莱克6051 小时前
C++ 七大排序算法详解:选择、冒泡、插入、计数、堆、快速与归并排序
数据结构·c++·算法·排序算法
shehuiyuelaiyuehao2 小时前
算法44,模拟算法,数青蛙
算法·哈希算法·散列表
Rocky Ding*2 小时前
一文读懂LLM Agent Skills 的运行时本质:从能力路由到渐进式披露
论文阅读·人工智能·深度学习·机器学习·aigc·ai-native·agent skills
hanlin032 小时前
刷题笔记:力扣第287题-寻找重复数
笔记·算法·leetcode
桃西西呀2 小时前
体检报告上一堆箭头看不懂?背后的逻辑就是随机森林:一群树投票,比一棵树靠谱
人工智能·机器学习·llm