MATLAB中3D二次判别分析(QDA)结果的曲面绘制方法问询
MATLAB中二次判别分析(QDA)判别曲面绘制方法
QDA判别曲面的正确方程
QDA的判别边界是两类后验概率相等的位置,推导后得到的二次方程为:
$$
\frac{1}{2}(x-\mu_1)T\Sigma_1{-1}(x-\mu_1) - \frac{1}{2}(x-\mu_2)T\Sigma_2{-1}(x-\mu_2) + \frac{1}{2}\ln\left(\frac{|\Sigma_1|}{|\Sigma_2|}\right) - \ln\left(\frac{\pi_1}{\pi_2}\right) = 0
$$
参数说明:
- $x = [x_1, x_2, x_3]^T$ 是你的3D特征向量
- $\mu_1、\mu_2$ 是两类的均值向量
- $\Sigma_1、\Sigma_2$ 是两类的协方差矩阵
- $\pi_1、\pi_2$ 是两类的先验概率(如果是等先验,最后一项直接为0)
对应到MATLAB的匿名函数(假设已用fitcdiscr得到QDA模型mdl),你可以这么写:
mu1 = mdl.Mu(1,:); mu2 = mdl.Mu(2,:); Sigma1 = mdl.Sigma(:,:,1); Sigma2 = mdl.Sigma(:,:,2); pi1 = mdl.Prior(1); pi2 = mdl.Prior(2); % 类1与类2的判别边界函数 qda_12 = @(x1,x2,x3) ... 0.5 * ([x1 x2 x3] - mu1) * inv(Sigma1) * ([x1; x2; x3] - mu1') ... - 0.5 * ([x1 x2 x3] - mu2) * inv(Sigma2) * ([x1; x2; x3] - mu2') ... + 0.5 * log(det(Sigma1)/det(Sigma2)) ... - log(pi1/pi2); % 类1与类3的判别边界函数 qda_13 = @(x1,x2,x3) ... 0.5 * ([x1 x2 x3] - mu1) * inv(Sigma1) * ([x1; x2; x3] - mu1') ... - 0.5 * ([x1 x2 x3] - mdl.Mu(3,:)) * inv(mdl.Sigma(:,:,3)) * ([x1; x2; x3] - mdl.Mu(3,:)') ... + 0.5 * log(det(Sigma1)/det(mdl.Sigma(:,:,3))) ... - log(pi1/mdl.Prior(3)); % 类2与类3的判别边界函数 qda_23 = @(x1,x2,x3) ... 0.5 * ([x1 x2 x3] - mu2) * inv(Sigma2) * ([x1; x2; x3] - mu2') ... - 0.5 * ([x1 x2 x3] - mdl.Mu(3,:)) * inv(mdl.Sigma(:,:,3)) * ([x1; x2; x3] - mdl.Mu(3,:)') ... + 0.5 * log(det(Sigma2)/det(mdl.Sigma(:,:,3))) ... - log(pi2/mdl.Prior(3));
把这三个函数分别传入你代码第135、137、139行的fimplicit3即可。
更省心的绘制方法
手动推导方程容易出错,推荐两种更简单的实现方式:
1. 网格点预测+等值面绘制
不需要推导任何方程,直接用模型预测网格点的类别,再画出类别间的边界:
% 确定数据范围,生成3D网格 xmin = min(X_data(:,1)); xmax = max(X_data(:,1)); ymin = min(X_data(:,2)); ymax = max(X_data(:,2)); zmin = min(X_data(:,3)); zmax = max(X_data(:,3)); [X1,X2,X3] = meshgrid(linspace(xmin,xmax,50), linspace(ymin,ymax,50), linspace(zmin,zmax,50)); X_grid = [X1(:), X2(:), X3(:)]; % 预测每个网格点的类别 Y_pred = predict(mdl, X_grid); Y_pred = reshape(Y_pred, size(X1)); % 绘制边界 figure; hold on; % 类1与类2的边界 isosurface(X1,X2,X3,Y_pred==1 & Y_pred~=2, 0.5); % 类1与类3的边界 isosurface(X1,X2,X3,Y_pred==1 & Y_pred~=3, 0.5); % 类2与类3的边界 isosurface(X1,X2,X3,Y_pred==2 & Y_pred~=3, 0.5); % 叠加原始数据点 scatter3(X_data(:,1), X_data(:,2), X_data(:,3), 10, Y_data, 'filled'); hold off;
这种方法对多类、高维场景都适用,完全不用关注方程形式。
2. 工具箱自带函数(限新版本)
如果你的MATLAB统计工具箱版本较新,plotDecisionBoundary函数可以直接传入QDA模型绘制边界,但该函数主要支持2D场景,3D下还是推荐上面的网格点方法。
注意事项
- 若协方差矩阵奇异,
inv(Sigma)会报错,可替换为pinv(Sigma)(伪逆),或用矩阵除法([x1;x2;x3]-mu1') \ Sigma1 \ ([x1 x2 x3]-mu1)计算二次型。 - 先验概率不要漏算,默认是样本的类别比例,遗漏会导致边界偏移。
内容的提问来源于stack exchange,提问作者user2587726
相关产品推荐
相关产品推荐

