使用scikit-learn中的KNN包实现对鸢尾花数据集的预测

1. 导入必要的库

首先,需要导入所需的库:

python 复制代码
import numpy as np
import pandas as pd
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import accuracy_score, classification_report
复制代码

2. 加载鸢尾花数据集

scikit-learn提供了方便的函数来加载鸢尾花数据集:

python 复制代码
# 加载鸢尾花数据集
iris = load_iris()
X = iris.data
y = iris.target

# 将数据集分为训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

3. 数据预处理

对数据进行标准化处理,以提高KNN算法的性能:

python 复制代码
# 标准化数据
scaler = StandardScaler()
X_train = scaler.fit_transform(X_train)
X_test = scaler.transform(X_test)

4. 训练KNN模型

使用KNeighborsClassifier来训练KNN模型:

python 复制代码
# 创建KNN分类器
knn = KNeighborsClassifier(n_neighbors=3)

# 训练模型
knn.fit(X_train, y_train)

5. 进行预测并评估模型

使用测试集进行预测,并评估模型的性能:

python 复制代码
# 进行预测
y_pred = knn.predict(X_test)

# 计算准确率
accuracy = accuracy_score(y_test, y_pred)
print(f'Accuracy: {accuracy:.2f}')

# 打印分类报告
print(classification_report(y_test, y_pred, target_names=iris.target_names))

6. 使用自定义数据集

如果有一个自定义的数据集,可以按照以下步骤进行操作。假设有一个CSV文件custom_dataset.csv,其中包含特征和标签。

python 复制代码
# 加载自定义数据集
custom_data = pd.read_csv('custom_dataset.csv')

# 假设特征列为前n-1列,标签列为最后一列
X_custom = custom_data.iloc[:, :-1].values
y_custom = custom_data.iloc[:, -1].values

# 将数据集分为训练集和测试集
X_custom_train, X_custom_test, y_custom_train, y_custom_test = train_test_split(X_custom, y_custom, test_size=0.2, random_state=42)

# 标准化数据
scaler_custom = StandardScaler()
X_custom_train = scaler_custom.fit_transform(X_custom_train)
X_custom_test = scaler_custom.transform(X_custom_test)

# 创建KNN分类器并训练
knn_custom = KNeighborsClassifier(n_neighbors=3)
knn_custom.fit(X_custom_train, y_custom_train)

# 进行预测
y_custom_pred = knn_custom.predict(X_custom_test)

# 计算准确率
accuracy_custom = accuracy_score(y_custom_test, y_custom_pred)
print(f'Custom Dataset Accuracy: {accuracy_custom:.2f}')

# 如果有标签名称,可以打印分类报告
# print(classification_report(y_custom_test, y_custom_pred, target_names=[...]))
相关推荐
醍醐实验室2 分钟前
量化推理实战:AWQ、GPTQ 与 SmoothQuant 激活异常值治理
人工智能
加速财经6 分钟前
2026年网站部署平台怎么选:6款平台,覆盖网页上线与 AI 项目部署
人工智能
长沙京卓9 分钟前
可源码交付|县域低空无人机智能运营平台:AI辅助飞行、巡检闭环与警用低空安防
人工智能·无人机
Python自动化直播10 分钟前
从脚本到语义:AI直播中控技术架构的三次演进与合规设计
人工智能·架构
pjj1985410 分钟前
深度学习-自定义Dataset和网络结构
人工智能
程序员cxuan15 分钟前
跟 WebUI 说再见了,最强 DeepSeek 桌面端来了!
人工智能·后端·程序员
计算机编程-吉哥19 分钟前
基于YOLO11s的苹果叶片病害检测系统 | 5类病害、2万+数据集、全栈闭环【计算机毕业设计选题推荐】
深度学习·毕业设计·课程设计·计算机毕业设计选题·机器学习毕业设计·大数据毕业设计选题推荐
haon112222 分钟前
树模型在信贷风控怎么用——从决策树到 XGBoost
大数据·人工智能·算法·决策树·机器学习·数据挖掘
wordbaby24 分钟前
BM25 是什么?手把手拆解搜索引擎的核心算法
人工智能·算法
古少侠24 分钟前
AI 内容粘到 Word 就乱?用 DS随心转 把复制的内容直接变成排版好的一段
人工智能·c#·word