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

如何在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')

代码说明

  1. 网格点生成:用linspace和meshgrid创建覆盖数据范围的X、Y网格,确保判别平面能完整覆盖数据区域。
  2. 平面计算:通过判别平面方程解出Z的表达式,代入网格点得到平面的Z坐标。
  3. 可视化优化:设置FaceAlpha让平面半透明,避免遮挡数据点;关闭EdgeColor减少视觉干扰。

内容的提问来源于stack exchange,提问作者user2587726

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 14:23:18