鸢尾花分类实战:从数据探索到模型部署全流程
1. 项目概述
鸢尾花分类是机器学习领域最经典的入门项目之一,相当于编程界的"Hello World"。这个项目使用Python的scikit-learn库(简称sklearn)来实现对鸢尾花数据集的分类任务。我最近在带学生做这个实验时发现,虽然网上教程很多,但大多数都停留在简单调用API的层面,缺乏对背后原理和实际工程考量的深入讲解。今天我就从一线教学和实践的角度,带大家完整走一遍这个项目的全流程,并分享一些只有实际做过才知道的细节技巧。
鸢尾花数据集包含3个品种(山鸢尾、变色鸢尾和维吉尼亚鸢尾)各50个样本,每个样本有4个特征(萼片长度、萼片宽度、花瓣长度、花瓣宽度)。我们的目标是构建一个机器学习模型,能够根据这4个特征准确预测花的品种。这个项目看似简单,却涵盖了机器学习从数据探索、预处理、模型训练到评估的完整流程,是理解监督学习分类任务的绝佳起点。
2. 环境准备与数据加载
2.1 开发环境配置
我推荐使用Anaconda创建独立的Python环境,避免包依赖冲突。以下是具体步骤:
conda create -n iris_classification python=3.8
conda activate iris_classification
pip install scikit-learn pandas matplotlib seaborn numpy
注意:虽然最新的Python版本是3.10+,但考虑到部分机器学习库的兼容性,建议使用3.8这个长期支持版本。我在实际教学中发现,3.9及以上版本有时会遇到sklearn的某些依赖项兼容问题。
2.2 数据加载与初步探索
sklearn内置了鸢尾花数据集,我们可以直接加载:
from sklearn.datasets import load_iris
import pandas as pd
iris = load_iris()
df = pd.DataFrame(iris.data, columns=iris.feature_names)
df['target'] = iris.target
df['species'] = df['target'].map({0:'setosa', 1:'versicolor', 2:'virginica'})
加载后建议立即进行数据探索:
print(df.head()) # 查看前5行数据
print(df.describe()) # 统计描述
print(df['species'].value_counts()) # 类别分布
我在实际项目中发现,很多初学者会忽略这个步骤直接开始建模,这是非常危险的。数据探索能帮助我们:
- 发现异常值(如花瓣长度出现负数)
- 了解特征量纲差异(萼片厘米级,花瓣毫米级)
- 确认类别是否平衡(本例中各50个样本)
3. 数据可视化与特征分析
3.1 单变量分布分析
使用seaborn绘制各特征的分布情况:
import seaborn as sns
import matplotlib.pyplot as plt
plt.figure(figsize=(12, 8))
for i, feature in enumerate(iris.feature_names):
plt.subplot(2, 2, i+1)
sns.histplot(data=df, x=feature, hue='species', kde=True)
plt.tight_layout()
plt.show()
这个可视化能直观展示:
- setosa的花瓣尺寸明显小于其他两类
- versicolor和virginica在花瓣特征上有部分重叠
- 所有特征都近似正态分布,没有极端离群值
3.2 特征相关性分析
计算并可视化特征间的Pearson相关系数:
plt.figure(figsize=(10, 8))
sns.heatmap(df[iris.feature_names].corr(), annot=True, cmap='coolwarm')
plt.title('Feature Correlation Matrix')
plt.show()
从我的经验看,高度相关的特征(如花瓣长度和宽度)可以考虑只保留一个,但在这个教学项目中我们保留全部特征,以便后续演示特征选择的影响。
4. 数据预处理
4.1 训练集测试集划分
from sklearn.model_selection import train_test_split
X = df[iris.feature_names]
y = df['target']
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42, stratify=y)
关键参数说明:
test_size=0.2:保留20%数据做最终测试random_state=42:固定随机种子确保结果可复现stratify=y:保持训练集和测试集的类别比例相同
实际工程中,我建议至少进行3次不同的随机划分来验证模型稳定性,教学演示为简化流程只做一次划分。
4.2 特征标准化
虽然决策树类算法不需要特征缩放,但为了演示完整流程,我们仍进行标准化:
from sklearn.preprocessing import StandardScaler
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test) # 注意:使用训练集的参数转换测试集
这里有个常见错误:在测试集上调用 fit_transform 会导致数据泄露。正确的做法是只在训练集上fit,然后统一transform两个数据集。
5. 模型训练与评估
5.1 逻辑回归模型
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import classification_report, confusion_matrix
lr = LogisticRegression(max_iter=200, multi_class='ovr')
lr.fit(X_train_scaled, y_train)
y_pred = lr.predict(X_test_scaled)
print("Classification Report:\n", classification_report(y_test, y_pred))
print("Confusion Matrix:\n", confusion_matrix(y_test, y_pred))
参数说明:
max_iter=200:增加迭代次数确保收敛multi_class='ovr':使用一对多策略处理多分类
5.2 支持向量机(SVM)
from sklearn.svm import SVC
svm = SVC(kernel='rbf', C=1.0, gamma='scale')
svm.fit(X_train_scaled, y_train)
y_pred = svm.predict(X_test_scaled)
print("SVM Classification Report:\n", classification_report(y_test, y_pred))
5.3 随机森林
from sklearn.ensemble import RandomForestClassifier
rf = RandomForestClassifier(n_estimators=100, max_depth=3, random_state=42)
rf.fit(X_train, y_train) # 决策树不需要特征缩放
y_pred = rf.predict(X_test)
print("Random Forest Classification Report:\n", classification_report(y_test, y_pred))
6. 模型比较与选择
将三个模型在测试集上的表现汇总:
| 模型 | 准确率 | 精确率(加权) | 召回率(加权) | F1分数(加权) |
|---|---|---|---|---|
| 逻辑回归 | 0.97 | 0.97 | 0.97 | 0.97 |
| SVM | 1.00 | 1.00 | 1.00 | 1.00 |
| 随机森林 | 0.93 | 0.94 | 0.93 | 0.93 |
从结果看,SVM表现最好,但要注意:
- 小数据集上可能存在偶然性
- SVM对参数更敏感(我们使用了默认参数)
- 随机森林没有做超参数调优
7. 模型优化与调参
7.1 网格搜索交叉验证
以SVM为例演示超参数调优:
from sklearn.model_selection import GridSearchCV
param_grid = {
'C': [0.1, 1, 10, 100],
'gamma': ['scale', 'auto', 0.1, 1],
'kernel': ['rbf', 'linear']
}
grid_search = GridSearchCV(SVC(), param_grid, cv=5, verbose=2)
grid_search.fit(X_train_scaled, y_train)
print("Best parameters:", grid_search.best_params_)
print("Best cross-validation score:", grid_search.best_score_)
7.2 学习曲线分析
检查模型是否过拟合或欠拟合:
from sklearn.model_selection import learning_curve
import numpy as np
train_sizes, train_scores, test_scores = learning_curve(
SVC(C=10, gamma=0.1, kernel='rbf'),
X_train_scaled, y_train, cv=5,
train_sizes=np.linspace(0.1, 1.0, 10)
)
plt.figure(figsize=(10, 6))
plt.plot(train_sizes, np.mean(train_scores, axis=1), 'o-', label="Training score")
plt.plot(train_sizes, np.mean(test_scores, axis=1), 'o-', label="Cross-validation score")
plt.xlabel("Training examples")
plt.ylabel("Score")
plt.legend()
plt.show()
8. 模型部署与应用
8.1 保存和加载模型
import joblib
# 保存最佳模型
joblib.dump(grid_search.best_estimator_, 'iris_svm_model.pkl')
# 加载模型
loaded_model = joblib.load('iris_svm_model.pkl')
# 使用模型预测新数据
new_data = [[5.1, 3.5, 1.4, 0.2]] # 示例数据
new_data_scaled = scaler.transform(new_data) # 使用相同的scaler
prediction = loaded_model.predict(new_data_scaled)
print("Predicted class:", iris.target_names[prediction][0])
8.2 构建简单的预测API
使用Flask创建Web服务:
from flask import Flask, request, jsonify
import joblib
import numpy as np
app = Flask(__name__)
model = joblib.load('iris_svm_model.pkl')
scaler = joblib.load('iris_scaler.pkl') # 需要单独保存scaler
@app.route('/predict', methods=['POST'])
def predict():
data = request.get_json()
features = [data['sepal_length'], data['sepal_width'],
data['petal_length'], data['petal_width']]
scaled_features = scaler.transform([features])
prediction = model.predict(scaled_features)
return jsonify({'species': iris.target_names[prediction[0]]})
if __name__ == '__main__':
app.run(host='0.0.0.0', port=5000)
9. 项目扩展与进阶方向
9.1 尝试其他分类算法
- K近邻(KNN)
- 朴素贝叶斯
- 梯度提升树(XGBoost, LightGBM)
9.2 特征工程进阶
- 尝试特征选择(如基于卡方检验、互信息)
- 创建新特征(如花瓣面积=长度×宽度)
- 使用PCA降维可视化
9.3 模型解释性
- 使用SHAP值解释模型预测
- 绘制决策边界
- 分析特征重要性
10. 常见问题与解决方案
10.1 模型准确率低
可能原因:
- 数据预处理不当(如未处理异常值)
- 特征间量纲差异大但未标准化
- 类别不平衡(虽然本数据集平衡)
解决方案:
- 重新检查数据质量
- 尝试不同的预处理方法
- 使用类别权重参数(如
class_weight='balanced')
10.2 过拟合问题
识别方法:
- 训练集准确率远高于验证集
- 学习曲线显示大间隙
解决方法:
- 增加正则化(如SVM的C参数)
- 获取更多数据
- 简化模型复杂度
10.3 预测结果不稳定
可能原因:
- 随机种子未固定
- 数据划分比例不合理
- 模型对输入变化敏感
解决方法:
- 固定所有random_state
- 使用交叉验证代替单次划分
- 尝试更鲁棒的算法(如随机森林)
11. 工程实践建议
- 版本控制 :使用git管理代码和数据,特别是预处理步骤和模型参数
- 实验记录 :详细记录每次实验的参数和结果,推荐使用MLflow或Weights & Biases
- 自动化测试 :为数据验证和模型评估编写单元测试
- 监控部署 :生产环境中监控模型性能衰减,建立回滚机制
我在实际教学中发现,学生最容易忽视的是第1和第2点。一个良好的实验记录习惯可以节省大量调试时间,特别是在尝试不同算法和参数组合时。建议为每个实验创建一个独立的Jupyter notebook或Python脚本,并添加详细的注释说明实验目的和观察结果。
更多推荐



所有评论(0)