如何在MATLAB中利用fitcdiscr输出绘制3D判别分析的判别平面?
3D线性判别分析(LDA)的判别平面可视化方案
核心思路
对于fitcdiscr训练出的线性判别分析模型,判别平面的方程由模型的**系数(Coefficients)和截距(Constant)**决定。两类之间的线性判别边界是平面,方程形式为:
$$w_1x + w_2y + w_3z + b = 0$$
其中$[w_1, w_2, w_3]$是判别函数的线性系数,$b$是截距项,这些参数均可从模型的Coefficients属性中提取。
关键参数提取
以3类分类场景为例,MdlLinear.Coefficients(i,j)对应第i类和第j类之间的判别函数:
MdlLinear.Coefficients(i,j).Linear:长度为3的向量,对应x、y、z三个特征的系数MdlLinear.Coefficients(i,j).Constant:判别函数的截距项
修改后的完整代码
clear clc close all %----------------------------生成模拟数据---------------------------- % 类别1数据 x1 = 1:0.01:3; r = -1 + 2*rand(1,201); x1 = x1 + r; y1 = 1:0.01:3; r = -1 + 2*rand(1,201); y1 = y1 + r; z1 = 1:0.01:3; r = -1 + 2*rand(1,201); z1 = z1 + r; x1 = x1'; y1 = y1'; z1 = z1'; label1 = ones(length(y1),1); Tclust1 = [label1,x1,y1,z1]; % 类别2数据 x2 = 4:0.01:6; r = -1 + 2*rand(1,201); x2 = x2 + r; y2 = 10:0.01:12; r = -1 + 2*rand(1,201); y2 = y2 + r; z2 = 10:0.01:12; r = -1 + 2*rand(1,201); z2 = z2 + r; x2 = x2'; y2 = y2'; z2 = z2'; label2 = 2*ones(length(y2),1); Tclust2 = [label2,x2,y2,z2]; % 类别3数据 x3 = 5:0.01:7; r = -1 + 2*rand(1,201); x3 = x3 + r; y3 = 11:0.01:13; r = -1 + 2*rand(1,201); y3 = y3 + r; z3 = 11:0.01:13; r = -1 + 2*rand(1,201); z3 = z3 + r; x3 = x3'; y3 = y3'; z3 = z3'; label3 = 3*ones(length(y3),1); Tclust3 = [label3,x3,y3,z3]; % 合并数据 T = [Tclust1;Tclust2;Tclust3]; data = T(:,2:4); labels = T(:,1); % 训练线性判别分析模型 MdlLinear = fitcdiscr(data,labels); %----------------------------可视化数据与判别平面---------------------------- figure() scatter3(x1,y1,z1,'g','filled') hold on scatter3(x2,y2,z2,'b','filled') scatter3(x3,y3,z3,'r','filled') title('3D线性判别分析数据与判别平面') xlabel('X特征') ylabel('Y特征') zlabel('Z特征') grid on axis equal % 生成网格点用于绘制平面 x_range = linspace(min(data(:,1)), max(data(:,1)), 50); y_range = linspace(min(data(:,2)), max(data(:,2)), 50); [X,Y] = meshgrid(x_range, y_range); % 绘制类1与类2的判别平面 coeff_12 = MdlLinear.Coefficients(1,2).Linear; const_12 = MdlLinear.Coefficients(1,2).Constant; % 从平面方程解出Z:coeff_12(1)*X + coeff_12(2)*Y + coeff_12(3)*Z + const_12 = 0 Z_12 = (-coeff_12(1)*X - coeff_12(2)*Y - const_12) / coeff_12(3); surf(X,Y,Z_12,'FaceColor','g','FaceAlpha',0.3,'EdgeColor','none') % 绘制类1与类3的判别平面 coeff_13 = MdlLinear.Coefficients(1,3).Linear; const_13 = MdlLinear.Coefficients(1,3).Constant; Z_13 = (-coeff_13(1)*X - coeff_13(2)*Y - const_13) / coeff_13(3); surf(X,Y,Z_13,'FaceColor','m','FaceAlpha',0.3,'EdgeColor','none') % 绘制类2与类3的判别平面 coeff_23 = MdlLinear.Coefficients(2,3).Linear; const_23 = MdlLinear.Coefficients(2,3).Constant; Z_23 = (-coeff_23(1)*X - coeff_23(2)*Y - const_23) / coeff_23(3); surf(X,Y,Z_23,'FaceColor','y','FaceAlpha',0.3,'EdgeColor','none') hold off legend('类别1','类别2','类别3','类1-类2边界','类1-类3边界','类2-类3边界','Location','best')
代码说明
- 网格点生成:用
linspace和meshgrid创建覆盖数据范围的X、Y网格,确保判别平面能完整覆盖数据区域。 - 平面计算:通过判别平面方程解出Z的表达式,代入网格点得到平面的Z坐标。
- 可视化优化:设置
FaceAlpha让平面半透明,避免遮挡数据点;关闭EdgeColor减少视觉干扰。
内容的提问来源于stack exchange,提问作者user2587726
相关产品推荐
相关产品推荐

