手写数字识别之绘制数字
1 课程概览
本课讲解如何通过索引读取CSV数据,将784个像素点reshape为28×28矩阵,并使用matplotlib绘制灰度图。详细讲解数据读取、索引越界判断、特征与标签分离、形状转换、图像绘制等完整流程,并使用Counter查看标签分布情况。
2 核心概念与定义
- iloc:pandas中基于位置索引获取数据的方法。
- values:获取Series的值(不含索引)。
- reshape:改变数组形状的方法,如reshape(28, 28)。
- imshow:matplotlib中显示图像的函数。
- cmap:颜色映射参数,'gray'表示灰度图。
- axis('off'):关闭坐标轴显示。
- Counter:collections模块中的计数器,用于统计标签分布。
3 算法与模型详解
3.1 绘制数字流程
- 读取CSV数据集
- 判断索引是否越界
- 分离特征(X)和标签(Y)
- 查看索引对应的数字
- 查看X的形状(784)
- 将784转换为28×28
- 绘制灰度图
3.2 数据读取与分离
| 操作 | 代码 | 说明 |
|---|---|---|
| 读取数据 | pd.read_csv() | 获取DataFrame |
| 获取特征 | df.iloc[:, 1:] | 所有行,第2列到最后 |
| 获取标签 | df.iloc[:, 0] | 所有行,第1列 |
3.3 索引越界判断
if index < 0 or index > len(df) - 1:
return # 索引越界,直接返回
说明:
- len(df) = 42001(含列名行)
- 有效索引范围:0 ~ 41999
- 索引越界时直接return
3.4 形状转换
| 转换前 | 转换后 | 代码 |
|---|---|---|
| (784,) | (28, 28) | reshape(28, 28) |
3.5 图像绘制参数
| 参数 | 说明 | 值 |
|---|---|---|
| X | 图像数据 | 28×28矩阵 |
| cmap | 颜色映射 | 'gray'(灰度图) |
| axis | 坐标轴 | 'off'(关闭) |
3.6 Counter统计标签分布
from collections import Counter
print(Counter(y)) # 统计各类别数量
作用:查看0~9各类别在数据集中的分布情况,判断数据是否均衡。
4 数学原理与推导
本课重点为代码实现,无数学推导。
5 代码示例
import pandas as pd
import matplotlib.pyplot as plt
from collections import Counter
def show_digit(index):
# 1. 读取数据集
df = pd.read_csv('./data/手写数字识别.csv')
print("数据形状:", df.shape) # (42000, 785)
# 2. 判断索引是否越界
if index < 0 or index > len(df) - 1:
print("索引越界")
return
# 3. 分离特征和标签
X = df.iloc[:, 1:] # 特征:所有行,第2列到最后(784列)
y = df.iloc[:, 0] # 标签:所有行,第1列
# 4. 查看索引对应的数字
print(f"该图片对应的数字是:{y.iloc[index]}")
# 5. 查看X的形状
print("X的形状:", X.iloc[index].shape) # (784,)
# 6. 将784转换为28×28
data = X.iloc[index].values.reshape(28, 28)
print("转换后形状:", data.shape) # (28, 28)
# 7. 绘制灰度图
plt.imshow(data, cmap='gray')
plt.axis('off') # 关闭坐标轴
plt.show()
# 8. 查看标签分布
print("标签分布:", Counter(y))
if __name__ == '__main__':
show_digit(9) # 索引9对应数字3
6 重难点与易错提醒
- ❗重点:特征用df.iloc[:, 1:]获取,标签用df.iloc[:, 0]获取。
- ❗重点:绘制图像前必须将784个像素点reshape为28×28。
- ❗重点:使用values获取Series的值,否则包含索引。
- ⚠️易错:索引越界判断时len(df)需要减1。
- ⚠️易错:忘记调用values,导致reshape报错。
- ⚠️易错:忘记axis('off'),图像会显示坐标轴。
- 💡深入理解:cmap='gray'表示灰度图,0=黑,255=白。
7 课堂问答精选
Q: 为什么需要使用values?
A: 因为df.iloc[index]返回的是Series对象,包含索引和值两部分。使用values只获取值部分(784个像素点数据),否则reshape会报错。
Q: 为什么索引9对应的是数字3?
A: 因为CSV第1行是列名,数据从第2行开始。索引0对应第2行,索引9对应第11行。从0开始数到第9个数据,对应的就是数字3。
Q: Counter的作用是什么?
A: Counter用于统计各类别在数据集中的数量分布。例如查看0~9各类别分别有多少条数据,判断数据是否均衡。如果数据均衡,模型训练效果会更好。
8 本课小结
- 绘制流程:读取数据→判断越界→分离特征标签→reshape转换→绘制灰度图。
- 特征获取:df.iloc[:, 1:](所有行,第2列到最后)。
- 标签获取:df.iloc[:, 0](所有行,第1列)。
- 形状转换:reshape(28, 28)将784个像素点转为28×28矩阵。
- 图像绘制:imshow(data, cmap='gray'),axis('off')关闭坐标轴。
- Counter统计标签分布,判断数据均衡性。
9 延伸思考与实践
- 实践:修改index值,观察不同索引对应的数字。
- 思考:为什么数据均衡对模型训练重要?
- 预习:如何训练KNN模型并保存。