这份代码使用 K-近邻算法 (KNN) 对经典的鸢尾花数据集进行分类,并演示如何进行预测。

from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.neighbors import KNeighborsClassifier
import numpy as np

if __name__ == '__main__':
    iris = load_iris()
    data = iris.get('data')
    target = iris.get('target')
    x_train, x_test, y_train, y_test = train_test_split(data, target, test_size=0.2, random_state=0)
    KNN = KNeighborsClassifier(n_neighbors=5)
    KNN.fit(x_train, y_train)
    train_score = KNN.score(x_train, y_train)
    test_score = KNN.score(x_test, y_test)
    print('模型的准确率:', test_score)
    X1 = np.array([[4.3, 3, 1.1, 0.1], [6.3, 2.3, 4.4, 1.3], [7.7, 2.6, 6.9, 2.3]])
    prediction = KNN.predict(X1)
    k = iris.get('target_names')[prediction]
    print('第一朵花的种类为:', k[0])
    print('第二朵花的种类为:', k[1])
    print('第二朵花的种类为:', k[2])

代码解析

  • 导入库

    • from sklearn.datasets import load_iris: 导入鸢尾花数据集。
    • from sklearn.model_selection import train_test_split: 导入用于将数据集划分为训练集和测试集的函数。
    • from sklearn.neighbors import KNeighborsClassifier: 导入 KNN 分类器。
    • import numpy as np: 导入 NumPy 库。
  • 加载数据

    • iris = load_iris(): 加载鸢尾花数据集。
    • data = iris.get('data'): 获取数据集的特征数据。
    • target = iris.get('target'): 获取数据集的标签数据。
  • 数据预处理

    • x_train, x_test, y_train, y_test = train_test_split(data, target, test_size=0.2, random_state=0): 将数据集划分为 80% 的训练集和 20% 的测试集,并使用随机种子 0 保证每次划分的结果相同。
  • 模型训练

    • KNN = KNeighborsClassifier(n_neighbors=5): 创建一个 KNN 分类器,其中 n_neighbors=5 表示使用 5 个最近邻。
    • KNN.fit(x_train, y_train): 使用训练集对 KNN 分类器进行训练。
  • 模型评估

    • train_score = KNN.score(x_train, y_train): 计算模型在训练集上的准确率。
    • test_score = KNN.score(x_test, y_test): 计算模型在测试集上的准确率。
    • print('模型的准确率:', test_score): 打印模型的测试集准确率。
  • 预测

    • X1 = np.array([[4.3, 3, 1.1, 0.1], [6.3, 2.3, 4.4, 1.3], [7.7, 2.6, 6.9, 2.3]]): 创建一个包含三个新样本的数组。
    • prediction = KNN.predict(X1): 使用训练好的 KNN 模型对新样本进行预测。
    • k = iris.get('target_names')[prediction]:将预测结果转换为类别名称。
  • 输出结果

    • print('第一朵花的种类为:', k[0])
    • print('第二朵花的种类为:', k[1])
    • print('第二朵花的种类为:', k[2])

这份代码展示了如何使用 KNN 算法进行分类并进行预测。你可以根据实际情况修改代码,例如调整 n_neighbors 值、使用不同的数据集或添加其他机器学习功能。

Python KNN 鸢尾花分类:代码解析及实战

原文地址: https://www.cveoy.top/t/topic/mzKn 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录