鸢尾花分类的四种经典方法

实现鸢尾花分类的四种经典方法:最小距离分类器、K近邻法、感知器和Fisher线性判别。

鸢尾花分类的四种经典方法

1. 数据准备与探索

%% 鸢尾花分类:四种方法比较
clear; close all; clc;

% 加载鸢尾花数据集
load fisheriris;
X = meas;       % 特征数据 (150x4)
Y = species;    % 类别标签

% 将类别标签转换为数值
[classes, ~, Y_num] = unique(Y);
num_classes = length(classes);
fprintf('鸢尾花数据集: %d个样本, %d个特征, %d个类别\n', ...
        size(X,1), size(X,2), num_classes);
fprintf('类别: %s, %s, %s\n', classes{:});

% 数据可视化 - 前两个特征
figure('Position', [100, 100, 1200, 800]);
subplot(2,3,1);
gscatter(X(:,1), X(:,2), Y);
xlabel('花萼长度 (cm)'); ylabel('花萼宽度 (cm)');
title('鸢尾花数据分布(前两个特征)');
legend(classes, 'Location', 'best');
grid on;

% 数据标准化
X_normalized = zscore(X);

% 数据集划分:70%训练,30%测试
rng(42); % 设置随机种子确保可重复性
train_ratio = 0.7;
num_samples = size(X,1);
num_train = round(train_ratio * num_samples);
idx = randperm(num_samples);
train_idx = idx(1:num_train);
test_idx = idx(num_train+1:end);

X_train = X_normalized(train_idx, :);
Y_train = Y_num(train_idx);
X_test = X_normalized(test_idx, :);
Y_test = Y_num(test_idx);

fprintf('训练集: %d个样本, 测试集: %d个样本\n', num_train, num_samples-num_train);

2. 最小距离分类器

%% 方法1: 最小距离分类器 (最近质心分类器)
function [Y_pred, centroids, accuracy] = minDistanceClassifier(X_train, Y_train, X_test, Y_test)
    % 最小距离分类器实现
    % 计算每个类别的质心,然后将测试样本分配到最近的质心类别
    
    unique_classes = unique(Y_train);
    num_classes = length(unique_classes);
    num_features = size(X_train, 2);
    
    % 计算每个类别的质心
    centroids = zeros(num_classes, num_features);
    for i = 1:num_classes
        class_mask = (Y_train == unique_classes(i));
        centroids(i, :) = mean(X_train(class_mask, :), 1);
    end
    
    % 对测试集进行预测
    num_test = size(X_test, 1);
    Y_pred = zeros(num_test, 1);
    distances = zeros(num_test, num_classes);
    
    for i = 1:num_test
        for j = 1:num_classes
            % 计算欧氏距离
            distances(i, j) = norm(X_test(i, :) - centroids(j, :));
        end
        [~, Y_pred(i)] = min(distances(i, :));
    end
    
    % 计算准确率
    accuracy = sum(Y_pred == Y_test) / num_test;
    
    fprintf('最小距离分类器准确率: %.2f%%\n', accuracy * 100);
end

% 运行最小距离分类器
[Y_pred_md, centroids, acc_md] = minDistanceClassifier(X_train, Y_train, X_test, Y_test);

3. K近邻法

%% 方法2: K近邻法
function [Y_pred, accuracy, neighbors] = kNNClassifier(X_train, Y_train, X_test, Y_test, k)
    % K近邻分类器实现
    
    if nargin < 5
        k = 5; % 默认K值
    end
    
    num_test = size(X_test, 1);
    num_train = size(X_train, 1);
    Y_pred = zeros(num_test, 1);
    neighbors = zeros(num_test, k);
    
    % 计算所有测试样本与训练样本的距离
    for i = 1:num_test
        distances = zeros(num_train, 1);
        for j = 1:num_train
            % 计算欧氏距离
            distances(j) = norm(X_test(i, :) - X_train(j, :));
        end
        
        % 找到最近的k个邻居
        [~, sorted_idx] = sort(distances);
        k_nearest = sorted_idx(1:k);
        neighbors(i, :) = k_nearest';
        
        % 多数投票
        k_labels = Y_train(k_nearest);
        Y_pred(i) = mode(k_labels);
    end
    
    % 计算准确率
    accuracy = sum(Y_pred == Y_test) / num_test;
    
    fprintf('K近邻分类器 (K=%d) 准确率: %.2f%%\n', k, accuracy * 100);
end

% 运行K近邻分类器,测试不同的K值
k_values = [1, 3, 5, 7];
acc_knn = zeros(size(k_values));

for i = 1:length(k_values)
    [Y_pred_knn, acc_knn(i), ~] = kNNClassifier(X_train, Y_train, X_test, Y_test, k_values(i));
end

% 选择最佳K值
[best_acc, best_idx] = max(acc_knn);
best_k = k_values(best_idx);
fprintf('最佳K值: %d, 准确率: %.2f%%\n', best_k, best_acc * 100);

4. 感知器

%% 方法3: 感知器
function [Y_pred, accuracy, weights, errors] = perceptronClassifier(X_train, Y_train, X_test, Y_test, learning_rate, max_epochs)
    % 感知器分类器实现(多类别)
    
    if nargin < 5
        learning_rate = 0.1;
    end
    if nargin < 6
        max_epochs = 100;
    end
    
    unique_classes = unique(Y_train);
    num_classes = length(unique_classes);
    num_features = size(X_train, 2);
    num_train = size(X_train, 1);
    
    % 添加偏置项
    X_train_bias = [X_train, ones(num_train, 1)];
    X_test_bias = [X_test, ones(size(X_test,1), 1)];
    
    % 初始化权重矩阵 (每个类别一组权重)
    weights = randn(num_classes, num_features + 1) * 0.1;
    
    % 训练感知器
    errors = zeros(max_epochs, 1);
    
    for epoch = 1:max_epochs
        epoch_errors = 0;
        
        for i = 1:num_train
            % 前向传播
            scores = weights * X_train_bias(i, :)';
            [~, predicted] = max(scores);
            
            % 获取真实类别的one-hot编码
            actual = Y_train(i);
            
            % 更新权重
            if predicted ~= actual
                % 减少预测类别的权重
                weights(predicted, :) = weights(predicted, :) - learning_rate * X_train_bias(i, :);
                % 增加真实类别的权重
                weights(actual, :) = weights(actual, :) + learning_rate * X_train_bias(i, :);
                
                epoch_errors = epoch_errors + 1;
            end
        end
        
        errors(epoch) = epoch_errors;
        
        % 早停:如果没有错误或错误很少
        if epoch_errors == 0
            fprintf('感知器在第 %d 轮收敛\n', epoch);
            errors = errors(1:epoch);
            break;
        end
    end
    
    % 在测试集上预测
    num_test = size(X_test_bias, 1);
    Y_pred = zeros(num_test, 1);
    
    for i = 1:num_test
        scores = weights * X_test_bias(i, :)';
        [~, Y_pred(i)] = max(scores);
    end
    
    % 计算准确率
    accuracy = sum(Y_pred == Y_test) / num_test;
    
    fprintf('感知器分类器准确率: %.2f%%\n', accuracy * 100);
end

% 运行感知器分类器
[Y_pred_perceptron, acc_perceptron, weights, errors] = perceptronClassifier(X_train, Y_train, X_test, Y_test);

5. Fisher线性判别

%% 方法4: Fisher线性判别
function [Y_pred, accuracy, W, projected_data] = fisherClassifier(X_train, Y_train, X_test, Y_test)
    % Fisher线性判别分析(多类别)
    
    unique_classes = unique(Y_train);
    num_classes = length(unique_classes);
    num_features = size(X_train, 2);
    num_test = size(X_test, 1);
    
    % 计算总体均值
    overall_mean = mean(X_train, 1);
    
    % 计算类间散度矩阵 Sb 和类内散度矩阵 Sw
    Sb = zeros(num_features);
    Sw = zeros(num_features);
    
    for i = 1:num_classes
        % 当前类别的样本
        class_mask = (Y_train == unique_classes(i));
        X_class = X_train(class_mask, :);
        n_class = sum(class_mask);
        
        % 类别均值
        class_mean = mean(X_class, 1);
        
        % 类间散度
        mean_diff = (class_mean - overall_mean)';
        Sb = Sb + n_class * (mean_diff * mean_diff');
        
        % 类内散度
        X_centered = X_class - class_mean;
        Sw = Sw + (X_centered' * X_centered);
    end
    
    % 解决广义特征值问题:Sb * W = lambda * Sw * W
    [eigenvectors, eigenvalues] = eig(Sb, Sw);
    eigenvalues = diag(eigenvalues);
    
    % 选择前C-1个最大特征值对应的特征向量
    [~, sorted_idx] = sort(eigenvalues, 'descend');
    W = eigenvectors(:, sorted_idx(1:num_classes-1));
    
    % 投影数据
    projected_train = X_train * W;
    projected_test = X_test * W;
    projected_data.train = projected_train;
    projected_data.test = projected_test;
    
    % 在投影空间中使用最小距离分类器
    [Y_pred, ~, accuracy] = minDistanceClassifier(projected_train, Y_train, projected_test, Y_test);
    
    fprintf('Fisher线性判别准确率: %.2f%%\n', accuracy * 100);
end

% 运行Fisher分类器
[Y_pred_fisher, acc_fisher, W, projected_data] = fisherClassifier(X_train, Y_train, X_test, Y_test);

6. 综合比较与可视化

%% 结果比较与可视化
figure('Position', [100, 100, 1400, 1000]);

% 1. 原始数据分布(使用前两个特征)
subplot(2,3,1);
gscatter(X(:,1), X(:,2), Y);
xlabel('花萼长度 (cm)'); ylabel('花萼宽度 (cm)');
title('原始数据分布');
legend(classes, 'Location', 'best');
grid on;

% 2. Fisher投影结果
subplot(2,3,2);
gscatter(projected_data.train(:,1), projected_data.train(:,2), Y_train);
xlabel('第一判别函数'); ylabel('第二判别函数');
title('Fisher判别投影(训练集)');
legend(classes, 'Location', 'best');
grid on;

% 3. 感知器训练误差
subplot(2,3,3);
plot(1:length(errors), errors, 'b-', 'LineWidth', 2);
xlabel('训练轮数'); ylabel('错误样本数');
title('感知器训练过程');
grid on;

% 4. K值对准确率的影响
subplot(2,3,4);
plot(k_values, acc_knn*100, 'ro-', 'LineWidth', 2, 'MarkerSize', 8);
xlabel('K值'); ylabel('准确率 (%)');
title('K近邻法中K值对准确率的影响');
grid on;

% 5. 混淆矩阵 - 最小距离分类器
subplot(2,3,5);
confusion_mat_md = confusionmat(Y_test, Y_pred_md);
confusionchart(confusion_mat_md, classes);
title('最小距离分类器混淆矩阵');

% 6. 方法性能比较
subplot(2,3,6);
methods = {'最小距离', 'K近邻', '感知器', 'Fisher'};
accuracies = [acc_md, best_acc, acc_perceptron, acc_fisher] * 100;

bar(accuracies, 'FaceColor', [0.2, 0.6, 0.8]);
set(gca, 'XTickLabel', methods);
ylabel('准确率 (%)');
title('四种分类方法性能比较');
grid on;

% 在柱状图上添加数值标签
for i = 1:length(accuracies)
    text(i, accuracies(i)+1, sprintf('%.1f%%', accuracies(i)), ...
        'HorizontalAlignment', 'center', 'FontWeight', 'bold');
end

% 输出详细比较结果
fprintf('\n=== 四种分类方法性能比较 ===\n');
fprintf('方法\t\t\t准确率\n');
fprintf('---------------------------------\n');
fprintf('最小距离分类器\t\t%.2f%%\n', acc_md*100);
fprintf('K近邻法 (K=%d)\t\t%.2f%%\n', best_k, best_acc*100);
fprintf('感知器\t\t\t%.2f%%\n', acc_perceptron*100);
fprintf('Fisher线性判别\t\t%.2f%%\n', acc_fisher*100);

% 计算每个类别的详细性能指标
fprintf('\n=== 各类别详细性能 ===\n');
for method_idx = 1:4
    switch method_idx
        case 1
            Y_pred_current = Y_pred_md;
            method_name = '最小距离';
        case 2
            [Y_pred_current, ~, ~] = kNNClassifier(X_train, Y_train, X_test, Y_test, best_k);
            method_name = sprintf('K近邻(K=%d)', best_k);
        case 3
            Y_pred_current = Y_pred_perceptron;
            method_name = '感知器';
        case 4
            Y_pred_current = Y_pred_fisher;
            method_name = 'Fisher';
    end
    
    fprintf('\n%s方法:\n', method_name);
    for class_idx = 1:num_classes
        class_mask = (Y_test == class_idx);
        class_accuracy = sum(Y_pred_current(class_mask) == class_idx) / sum(class_mask);
        fprintf('  %s: %.2f%%\n', classes{class_idx}, class_accuracy*100);
    end
end

7. 扩展分析:特征重要性

%% 特征重要性分析
feature_names = {'花萼长度', '花萼宽度', '花瓣长度', '花瓣宽度'};

figure('Position', [100, 100, 1200, 400]);

% Fisher判别中的特征权重
subplot(1,2,1);
feature_weights = sum(abs(W), 2); % 特征在判别函数中的总权重
[~, feature_rank] = sort(feature_weights, 'descend');

bar(feature_weights, 'FaceColor', [0.8, 0.4, 0.2]);
set(gca, 'XTickLabel', feature_names);
ylabel('特征权重绝对值之和');
title('Fisher判别中的特征重要性');
grid on;

% 添加数值标签
for i = 1:length(feature_weights)
    text(i, feature_weights(i)+0.01, sprintf('%.3f', feature_weights(i)), ...
        'HorizontalAlignment', 'center');
end

% 仅使用前两个最佳特征的性能
subplot(1,2,2);
best_features = feature_rank(1:2);
X_train_reduced = X_train(:, best_features);
X_test_reduced = X_test(:, best_features);

% 使用减少的特征重新运行分类器
[~, acc_md_reduced] = minDistanceClassifier(X_train_reduced, Y_train, X_test_reduced, Y_test);
[~, acc_knn_reduced, ~] = kNNClassifier(X_train_reduced, Y_train, X_test_reduced, Y_test, best_k);

fprintf('\n=== 特征选择分析 ===\n');
fprintf('最重要的两个特征: %s, %s\n', feature_names{best_features});
fprintf('使用两个特征的性能:\n');
fprintf('  最小距离分类器: %.2f%% (原: %.2f%%)\n', acc_md_reduced*100, acc_md*100);
fprintf('  K近邻法: %.2f%% (原: %.2f%%)\n', acc_knn_reduced*100, best_acc*100);

方法原理总结

1. 最小距离分类器

2. K近邻法

3. 感知器

4. Fisher线性判别

参考代码 使用四种方法进行鸢尾花分类:最小距离分类器,K 近邻法,感知器,Fisher 准则。 www.youwenfan.com/contentzhf/65404.html

实际应用建议

  1. 数据量小且特征明显 → 最小距离分类器
  2. 数据分布复杂 → K近邻法(需调优K值)
  3. 需要在线学习 → 感知器
  4. 特征降维与可视化 → Fisher判别
  5. 实际项目 → 建议尝试更先进的算法(SVM、随机森林、神经网络)

 

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