You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.20 07:52:49