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()

更多推荐