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())  # 类别分布

我在实际项目中发现,很多初学者会忽略这个步骤直接开始建模,这是非常危险的。数据探索能帮助我们:

  1. 发现异常值(如花瓣长度出现负数)
  2. 了解特征量纲差异(萼片厘米级,花瓣毫米级)
  3. 确认类别是否平衡(本例中各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表现最好,但要注意:

  1. 小数据集上可能存在偶然性
  2. SVM对参数更敏感(我们使用了默认参数)
  3. 随机森林没有做超参数调优

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 模型准确率低

可能原因:

  1. 数据预处理不当(如未处理异常值)
  2. 特征间量纲差异大但未标准化
  3. 类别不平衡(虽然本数据集平衡)

解决方案:

  • 重新检查数据质量
  • 尝试不同的预处理方法
  • 使用类别权重参数(如 class_weight='balanced'

10.2 过拟合问题

识别方法:

  • 训练集准确率远高于验证集
  • 学习曲线显示大间隙

解决方法:

  • 增加正则化(如SVM的C参数)
  • 获取更多数据
  • 简化模型复杂度

10.3 预测结果不稳定

可能原因:

  1. 随机种子未固定
  2. 数据划分比例不合理
  3. 模型对输入变化敏感

解决方法:

  • 固定所有random_state
  • 使用交叉验证代替单次划分
  • 尝试更鲁棒的算法(如随机森林)

11. 工程实践建议

  1. 版本控制 :使用git管理代码和数据,特别是预处理步骤和模型参数
  2. 实验记录 :详细记录每次实验的参数和结果,推荐使用MLflow或Weights & Biases
  3. 自动化测试 :为数据验证和模型评估编写单元测试
  4. 监控部署 :生产环境中监控模型性能衰减,建立回滚机制

我在实际教学中发现,学生最容易忽视的是第1和第2点。一个良好的实验记录习惯可以节省大量调试时间,特别是在尝试不同算法和参数组合时。建议为每个实验创建一个独立的Jupyter notebook或Python脚本,并添加详细的注释说明实验目的和观察结果。

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐