基于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);
五、扩展应用场景
-
手写数字识别
% 加载MNIST数据集 [XTrain, YTrain] = digitTrain4DArrayData; XTest = digitTest4DArrayData; % 特征提取(展平图像) XTrain = reshape(XTrain, [], 28 * 28); XTest = reshape(XTest, [], 28 * 28); -
生物信息学分类
% 基因表达数据分析 load('gene_expression.mat'); X = normalize(geneData); Y = disease_labels;