传统机器学习分类模型详解(一)
·
- 逻辑回归
- 支持向量机
- 决策树
概述
传统机器学习分类模型是指在深度学习兴起之前广泛使用的分类算法,它们具有原理清晰、计算效率高、可解释性强等特点。
主要分类模型
- 逻辑回归
原理介绍
逻辑回归是一种线性分类模型,通过Sigmoid函数将线性输出映射到概率空间。
核心公式:
P(y=1|x) = 1 / (1 + e^(-wᵀx))
代码实现
import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, classification_report
from sklearn.preprocessing import StandardScaler
# 示例:使用鸢尾花数据集
from sklearn.datasets import load_iris
# 加载数据
iris = load_iris()
X, y = iris.data, iris.target
# 只使用两类进行二分类示例
X_binary = X[y != 2]
y_binary = y[y != 2]
# 数据标准化
scaler = StandardScaler()
X_scaled = scaler.fit_transform(X_binary)
# 划分训练测试集
X_train, X_test, y_train, y_test = train_test_split(
X_scaled, y_binary, test_size=0.3, random_state=42
)
# 创建并训练模型
log_reg = LogisticRegression(random_state=42)
log_reg.fit(X_train, y_train)
# 预测
y_pred = log_reg.predict(X_test)
# 评估
print("准确率:", accuracy_score(y_test, y_pred))
print("分类报告:")
print(classification_report(y_test, y_pred))
# 查看系数(可解释性)
print("模型系数:", log_reg.coef_)
print("截距:", log_reg.intercept_)
优缺点
优点:
计算效率高,训练速度快
输出具有概率意义
可解释性强,系数代表特征重要性
不容易过拟合
缺点:
只能处理线性可分问题
对特征相关性敏感
需要特征工程
- 支持向量机
原理介绍
SVM通过寻找最大间隔超平面来实现分类,可以使用核函数处理非线性问题。
代码实现
from sklearn.svm import SVC
from sklearn.metrics import confusion_matrix
import matplotlib.pyplot as plt
import seaborn as sns
# 创建SVM模型
svm_model = SVC(kernel='rbf', C=1.0, gamma='scale', random_state=42)
svm_model.fit(X_train, y_train)
# 预测
y_pred_svm = svm_model.predict(X_test)
print("SVM准确率:", accuracy_score(y_test, y_pred_svm))
# 可视化混淆矩阵
cm = confusion_matrix(y_test, y_pred_svm)
plt.figure(figsize=(8, 6))
sns.heatmap(cm, annot=True, fmt='d', cmap='Blues')
plt.title('SVM混淆矩阵')
plt.ylabel('真实标签')
plt.xlabel('预测标签')
plt.show()
优缺点
优点:
在高维空间表现良好
核技巧可以处理非线性问题
泛化能力强
缺点:
大规模数据训练慢
参数调优复杂
结果可解释性差
- 决策树
原理介绍
通过树形结构进行决策,基于信息增益或基尼不纯度选择最佳分割特征。
代码实现
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.model_selection import cross_val_score
# 创建决策树模型
dt_model = DecisionTreeClassifier(
max_depth=3,
min_samples_split=2,
min_samples_leaf=1,
random_state=42
)
dt_model.fit(X_train, y_train)
# 预测
y_pred_dt = dt_model.predict(X_test)
print("决策树准确率:", accuracy_score(y_test, y_pred_dt))
# 可视化决策树
plt.figure(figsize=(12, 8))
plot_tree(dt_model,
feature_names=iris.feature_names,
class_names=iris.target_names[:2],
filled=True,
rounded=True)
plt.title('决策树可视化')
plt.show()
# 特征重要性
feature_importance = dt_model.feature_importances_
plt.barh(iris.feature_names, feature_importance)
plt.title('特征重要性')
plt.show()
优缺点
优点:
直观易懂,可解释性强
不需要特征缩放
能够处理数值和类别特征
缺点:
容易过拟合
对数据微小变化敏感
可能创建 biased 树
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐


所有评论(0)