如何用Python基于CSV图像数据生成SVM与朴素贝叶斯的准确率及混淆矩阵
Python实现SVM与朴素贝叶斯分类(CSV图像数据)
一、整体思路概述
咱们的需求是处理CSV格式的图像数据(通常每行对应一个样本,第一列是类别标签,后续列是像素特征值),用Python实现SVM和朴素贝叶斯分类,最后输出两类模型的准确率和混淆矩阵。核心依赖scikit-learn库,搭配pandas做数据加载、numpy做数值处理,全程都是Python生态里的常用工具,上手很顺。
二、步骤拆解与代码实现
1. 依赖库安装
先确保你装好了所需工具包,终端执行:
pip install pandas numpy scikit-learn
2. 数据加载与预处理
CSV图像数据的常规结构:第一列是label(比如数字0-9、猫/狗分类标签),后面的列是每个像素的灰度值(0-255)或RGB通道值。
import pandas as pd import numpy as np from sklearn.model_selection import train_test_split from sklearn.preprocessing import StandardScaler # 1. 加载CSV数据(替换成你的文件路径) data = pd.read_csv("your_image_data.csv") # 2. 拆分特征与标签 X = data.drop("label", axis=1) # 特征:所有像素列 y = data["label"] # 标签:类别列 # 3. 划分训练集和测试集(8:2比例,stratify保证类别分布一致) X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, random_state=42, stratify=y ) # 4. 特征归一化(SVM对特征尺度敏感,必须做;朴素贝叶斯可省略,但做了也不影响) scaler = StandardScaler() X_train_scaled = scaler.fit_transform(X_train) X_test_scaled = scaler.transform(X_test)
3. SVM模型训练与评估
from sklearn.svm import SVC from sklearn.metrics import accuracy_score, confusion_matrix # 初始化SVM分类器(默认RBF核,高维图像数据也可以试试线性核linear,速度更快) svm_clf = SVC(kernel="rbf", random_state=42) # 训练模型 svm_clf.fit(X_train_scaled, y_train) # 预测测试集 y_pred_svm = svm_clf.predict(X_test_scaled) # 计算准确率 svm_acc = accuracy_score(y_test, y_pred_svm) print(f"SVM分类准确率: {svm_acc:.4f}") # 生成混淆矩阵 svm_cm = confusion_matrix(y_test, y_pred_svm) print("SVM混淆矩阵:") print(svm_cm)
4. 朴素贝叶斯模型训练与评估
图像像素是连续值(0-255),咱们用高斯朴素贝叶斯(GaussianNB)最合适;如果是离散计数特征可以换MultinomialNB:
from sklearn.naive_bayes import GaussianNB # 初始化高斯朴素贝叶斯分类器 nb_clf = GaussianNB() # 训练模型(朴素贝叶斯不需要归一化,用原始特征也可以,这里用归一化后的也没问题) nb_clf.fit(X_train_scaled, y_train) # 预测测试集 y_pred_nb = nb_clf.predict(X_test_scaled) # 计算准确率 nb_acc = accuracy_score(y_test, y_pred_nb) print(f"朴素贝叶斯分类准确率: {nb_acc:.4f}") # 生成混淆矩阵 nb_cm = confusion_matrix(y_test, y_pred_nb) print("朴素贝叶斯混淆矩阵:") print(nb_cm)
5. 可选:可视化混淆矩阵
如果想更直观地看混淆矩阵,用seaborn画热力图就行:
pip install seaborn matplotlib
import seaborn as sns import matplotlib.pyplot as plt # 定义绘图函数 def plot_confusion_matrix(cm, class_names, title): plt.figure(figsize=(8,6)) sns.heatmap(cm, annot=True, fmt="d", cmap="Blues", xticklabels=class_names, yticklabels=class_names) plt.title(title) plt.xlabel("预测标签") plt.ylabel("真实标签") plt.show() # 获取类别名称(假设你的标签是0-9,替换成实际类别名即可) class_names = np.unique(y).astype(str) # 绘制SVM混淆矩阵 plot_confusion_matrix(svm_cm, class_names, "SVM分类混淆矩阵") # 绘制朴素贝叶斯混淆矩阵 plot_confusion_matrix(nb_cm, class_names, "朴素贝叶斯分类混淆矩阵")
三、注意事项
- 如果你的CSV没有表头,记得用
pd.read_csv("your_data.csv", header=None),然后手动指定标签列:比如y = data.iloc[:, 0],X = data.iloc[:, 1:]。 - SVM核函数选择:线性核适合高维图像数据,训练速度快;RBF核适合非线性分布,但大数据集下训练时间会更长,按需调整。
- 朴素贝叶斯假设特征独立,对于图像这种特征有相关性的数据,效果通常不如SVM,但胜在训练速度极快,适合做快速基准模型。
内容的提问来源于stack exchange,提问作者MASYAM
相关产品推荐
相关产品推荐

