基于MATLAB的SVM多分类实现方案

基于MATLAB的SVM多分类实现方案

一、代码

%% 1. 数据加载与预处理
% 加载鸢尾花数据集(4类多分类问题)
load fisheriris
X = meas(:,1:2); % 选取前两个特征便于可视化
Y = species;

% 数据标准化(重要!SVM对尺度敏感)
[X_scaled, mu, sigma] = zscore(X);

%% 2. 数据集划分(训练集70%+测试集30%)
cv = cvpartition(Y,'HoldOut',0.3);
X_train = X_scaled(cv.training,:);
Y_train = Y(cv.training);
X_test = X_scaled(cv.test,:);
Y_test = Y(cv.test);

%% 3. 多分类SVM模型构建
% 定义多分类策略(OvR或OvO)
t = templateSVM('KernelFunction','rbf','Standardize',false);

% 使用ECOC编码实现多分类
model = fitcecoc(X_train, Y_train, ...
    'Learners', t, ...
    'Coding', 'onevsall', % 可选 'onevsone'
    'CrossVal', 'on', ...
    'OptimizeHyperparameters', 'auto');

%% 4. 模型训练与预测
Y_pred = predict(model, X_test);

%% 5. 性能评估
% 混淆矩阵
C = confusionmat(Y_test, Y_pred);
disp('混淆矩阵:');
disp(C);

% 分类报告
disp('分类报告:');
disp(classificationReport(Y_test, Y_pred));

% ROC曲线(多类需转换为二分类计算)
figure;
plotroc(Y_test, model.predictorScores);

%% 6. 可视化决策边界(仅限二维特征)
figure;
gscatter(X_scaled(:,1), X_scaled(:,2), Y);
hold on;
h = 0.02; % 网格步长
[x1Grid,x2Grid] = meshgrid(min(X_scaled(:,1)):h:max(X_scaled(:,1)),...
    min(X_scaled(:,2)):h:max(X_scaled(:,2)));
scores = predict(model, [x1Grid(:), x2Grid(:)]);
contourf(x1Grid, x2Grid, reshape(scores(:,2),size(x1Grid)),...
    'LineStyle','none', 'Alpha',0.3);
title('SVM多分类决策边界');
hold off;

二、关键参数说明与优化

1. 多分类策略对比

策略类型 实现方式 适用场景 计算复杂度
OvR 为每个类训练一个二分类器 类别数适中(K≤10) O(K)
OvO 每对类别训练一个分类器 类别数较少(K≤5) O(K²)

2. 核函数选择指南

% 常用核函数配置
model = fitcsvm(X_train, Y_train, ...
    'KernelFunction', 'rbf',...  % 径向基核(默认)
    'KernelScale', 'auto',...    % 自动调整核参数
    'BoxConstraint', 1);        % 正则化参数C

3. 参数调优方法

% 网格搜索优化(示例:C和gamma优化)
tuneGrid = struct('BoxConstraint',[0.1,1,10],...
                  'KernelScale',[0.5,1,2]);

optimizedModel = fitcsvm(X_train, Y_train, ...
    'KernelFunction','rbf',...
    'OptimizeHyperparameters','grid',...
    'HyperparameterOptimizationOptions',...
    struct('AcquisitionFunctionName','expected-improvement-plus'));

三、性能评估指标

1. 分类指标计算

% 计算准确率
accuracy = sum(Y_pred == Y_test)/numel(Y_test);

% 计算混淆矩阵
C = confusionmat(Y_test, Y_pred);

% 计算精确率、召回率、F1-score
precision = diag(C)./sum(C,2);
recall = diag(C)./sum(C,1)';
f1_score = 2*(precision.*recall)./(precision+recall);

2. ROC曲线分析

% 多类ROC分析(需转换为二分类问题)
model = fitcecoc(X_train, Y_train, 'Coding', 'onevsall');
[~,~,~,AUC] = perfcurve(Y_test, model.predictorScores(:,2), 'versicolor');
disp(['AUC值: ', num2str(AUC)]);

参考代码 支持向量机,用于分类,含训练集与测试集,用SVM进行多分类 www.youwenfan.com/contentcsi/64920.html

四、优化建议

1. 大规模数据处理

% 使用线性SVM加速
model = fitcsvm(X_train, Y_train, 'KernelFunction','linear', 'Dual',false);

% 增量学习(适用于流式数据)
model = incrementalClassificationSVM();
model = incrementalLearner(model);

2. 类别不平衡处理

% 设置类别权重
model = fitcsvm(X_train, Y_train, ...
    'ClassNames',{'setosa','versicolor','virginica'},...
    'Prior',[0.5,0.3,0.2]);

3. 特征工程优化

% 特征选择(使用递归特征消除)
[rfeModel, selectedFeatures] = rfe(X_train, Y_train, 2);

% 特征变换(PCA降维)
[coeff,score,latent] = pca(X_train);
X_train_pca = score(:,1:2);

五、扩展应用场景

  1. 手写数字识别

    % 加载MNIST数据集
    [XTrain, YTrain] = digitTrain4DArrayData;
    XTest = digitTest4DArrayData;
    
    % 特征提取(展平图像)
    XTrain = reshape(XTrain, [], 28 * 28);
    XTest = reshape(XTest, [], 28 * 28);
    
  2. 生物信息学分类

    % 基因表达数据分析
    load('gene_expression.mat');
    X = normalize(geneData);
    Y = disease_labels;
    

 

专注于matlab/simulink,电子电路,编程