MATLAB实现线性判别分析(LDA)进行数据降维

MATLAB实现线性判别分析(LDA)进行数据降维


一、LDA降维原理

1. 数学模型

2. 降维维度限制


二、MATLAB实现方案

1. 基础实现(使用内置函数)

% 加载示例数据(鸢尾花数据集)
load fisheriris
X = meas;
Y = species;

% 执行LDA降维
ldaModel = fitcdiscr(X, Y, 'DiscrimType', 'linear');
X_lda = predict(ldaModel, X);

% 可视化结果
figure;
gscatter(X_lda(:,1), X_lda(:,2), Y);
xlabel('LD1');
ylabel('LD2');
title('LDA降维结果(内置函数)');

2. 手动实现(完整流程)

function [W, X_lda] = manual_LDA(X, Y, num_components)
    % 输入参数:
    % X: 原始数据矩阵 (n_samples x n_features)
    % Y: 类别标签 (n_samples x 1)
    % num_components: 降维后维度
    
    % 计算类别信息
    classes = unique(Y);
    n_classes = length(classes);
    n_features = size(X, 2);
    
    % 计算全局均值
    mu = mean(X);
    
    % 初始化散度矩阵
    Sw = zeros(n_features, n_features);
    Sb = zeros(n_features, n_features);
    
    % 计算类内/类间散度矩阵
    for i = 1:n_classes
        Xi = X(Y == classes(i), :);
        Ni = size(Xi, 1);
        mu_i = mean(Xi);
        
        % 类内散度矩阵
        Sw = Sw + (Xi - mu_i)' * (Xi - mu_i);
        
        % 类间散度矩阵
        Sb = Sb + Ni * (mu_i - mu)' * (mu_i - mu);
    end
    
    % 求解广义特征值问题
    [V, D] = eig(Sb, Sw);
    
    % 特征值排序
    [~, idx] = sort(diag(D), 'descend');
    V = V(:, idx);
    
    % 选择前num_components个特征向量
    W = V(:, 1:num_components);
    
    % 数据投影
    X_lda = X * W;
end

3. 调用示例

% 生成示例数据
rng(0);
X = [randn(50,2)+1; randn(50,2)-1];
Y = [ones(50,1); 2*ones(50,1)];

% 执行手动LDA
[X_lda, W] = manual_LDA(X, Y, 1);

% 可视化
figure;
scatter(X_lda(Y==1), zeros(sum(Y==1),1), 'r', 'filled');
hold on;
scatter(X_lda(Y==2), zeros(sum(Y==2),1), 'b', 'filled');
xlabel('投影值');
title('手动实现LDA降维');

三、应用场景

1. 人脸识别(ORL数据库)

% 加载ORL人脸数据集
[X, labels] = load_orl_database();

% LDA降维
X_lda = manual_LDA(X, labels, 2);

% 分类准确率验证
cv = cvpartition(labels, 'KFold', 5);
accuracy = crossval(@(Xtrain,Ytrain,Xtest,Ytest) ...
    sum(predict(manual_LDA(Xtrain,Ytrain), Xtest) == Ytest)/numel(Ytest), ...
    X, labels, 'partition', cv);
disp(['平均分类准确率: ', num2str(accuracy*100), '%']);

2. 高光谱图像分类

% 加载Indian Pines数据集
[X, ground_truth] = load_indian_pines();

% LDA降维预处理
X_lda = manual_LDA(X, ground_truth, 10);

% SVM分类
svmModel = fitcsvm(X_lda, ground_truth);
predicted = predict(svmModel, X_lda);
accuracy = sum(predicted == ground_truth)/numel(ground_truth);
disp(['分类准确率: ', num2str(accuracy*100), '%']);

四、常见问题解决

  1. 维度不足错误

    % 当类别数不足时自动调整维度
    num_components = min(num_components, length(classes)-1);
    
  2. 非数值型标签处理

    % 将类别标签转换为数值型
    Y = grp2idx(Y);
    
  3. 数据不平衡处理

    % 添加样本权重
    model = fitcdiscr(X, Y, 'DiscrimType', 'linear', 'Prior', 'uniform');
    

参考代码 利用LDA算法,实现数据降维 www.youwenfan.com/contentcsk/79531.html

五、工具包对比

工具特性 内置fitcdiscr 手动实现 优化版本
开发效率
计算性能 中等
功能扩展性 有限 完全可控 高度可定制
适用场景 快速验证 算法研究 工业级应用

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