Python KNN 鸢尾花分类:代码解析及实战
这份代码使用 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 值、使用不同的数据集或添加其他机器学习功能。
原文地址: https://www.cveoy.top/t/topic/mzKn 著作权归作者所有。请勿转载和采集!