机器学习案例:手写数字识别
·
import matplotlib.pyplot as plt
import pandas as pd
from sklearn.metrics import accuracy_score
from sklearn.model_selection import train_test_split
from sklearn.neighbors import KNeighborsClassifier
import joblib
from collections import Counter # 保存模型
import warnings
# 参一:忽略警告,参二,忽略模块
warnings.filterwarnings('ignore',module='sklearn')
# 定义函数,展示图片
def show_digit(idx):
df = pd.read_csv('./data/手写数字识别.csv')
print(df,'\n') # (42000 x 785)
if idx < 0 or idx > len(df) - 1:
print("索引超出范围")
return
x = df.iloc[:, 1:]# 取出所有行,从第二列开始的所有列作为特征
y = df.iloc[:, 0]# 只要取出第一列即可
# 查看用户传入的索引对应的图片是多少
print(f'索引{idx}对应的标签是:{y[idx]}')
print(f'查看所有标签的分布情况:{Counter(y)}')
# 看看有多少个0,多少个1,多少个...
# 将748*1转为28*28的格式
x=x.iloc[idx].values.reshape(28, 28)
# 具体绘制灰度图
plt.imshow(x,cmap='gray')
plt.axis('off')
plt.show()
pass
# 保存模型
def train_model():
# 数据加载
df = pd.read_csv('./data/手写数字识别.csv')
# 数据预处理
x = df.iloc[:,1:] #拆分特征列
y = df.iloc[:,0] #拆分标签列
#打印特征列和标签列的形状
print(f'x的形状:{x.shape}')
print(f'y的形状:{y.shape}')
# 数据归一化
print(x)
x = x/255 #(当前值-最小值)/(最大值-最小值)
# 拆分测试集和训练集
# 参一:特征列,参二:标签列,参三:测试集比例,参四:随机种子,参五,参考y值进行抽取,保持数据均衡
x_train,x_test,y_train,y_test = train_test_split(x,y,test_size=0.2,random_state=23,stratify=y)
# 模型训练
# 创建模型对象
estimator = KNeighborsClassifier(n_neighbors=3)
# 模型训练
estimator.fit(x_train,y_train)
# 模型评估
print(f'准确率{estimator.score(x_test,y_test)}')
print(f'准确率{accuracy_score(y_test,estimator.predict(x_test))}')
# 模型保存
joblib.dump(estimator, './My_model/knn手写数字识别.pkl')
print('模型保存成功')
# 使用模型
def use_model():
# 1,加载图片
x = plt.imread('./data/demo.png')
# 2,绘制图片
# plt.imshow(x, cmap='gray')
# plt.axis('off')
# plt.show()
# 加载模型
estimator = joblib.load('./My_model/knn手写数字识别.pkl')
# 模型预测
print(x.shape) #原先是28*28,现在转为1*784,一行784列
print(x.reshape(1,784).shape)
print(x.reshape(1, -1).shape) # 能转多少列就转多少列
# 查看数据集转换,此时的图还是原始数据集,未经过改造,也要归一化,因为训练模型时使用了归一化操作
# 但是读图时像素值并非精准,除以255可能会导致预测失败,除非能保证读到的数字序列和原始的数字序列都是
x = x.reshape(1,-1)
# 模型预测
y_pred = estimator.predict(x)
print(f'预测值为{y_pred}')
pass
# 测试
if __name__ == "__main__":
# 绘制数字
# show_digit(0)
# 训练模型并保存模型
# train_model()
# 模型预测(使用模型)
use_model()
更多推荐
所有评论(0)