基于Matlab实现双月分类的直线划分技术咨询及代码解析
如何用直线实现双月数据集的分类(基于Matlab)
首先,我先把你提供的双月数据生成代码补全完整,这样我们能得到标准的双月数据集,然后再一步步实现直线分类的方法。
第一步:补全双月数据生成代码
你的代码只生成了半个月亮,我补充了第二个月亮的生成逻辑,同时添加了标签和可视化:
function [data, labels] = dm(r, w, ts, d) clear all; close all; % 设置默认参数 if nargin < 4, d = -4; end if nargin < 3, ts = 1000; end if nargin < 2, w = 6; end if nargin < 1, r = 10; end ts1 = 10 * ts; done = 0; tmp1 = []; % 生成第一个月亮(上半部分) while ~done tmp = [2*(r + w/2)*(rand(ts1,1)-0.5) (r + w/2)*rand(ts1,1)]; tmp(:,3) = sqrt(tmp(:,1).^2 + tmp(:,2).^2); idx = find(tmp(:,3) > r - w/2 & tmp(:,3) < r + w/2); tmp1 = [tmp1; tmp(idx,1:2)]; if length(idx) >= ts done = 1; end end moon1 = tmp1(1:ts,:); % 生成第二个月亮(下半部分,偏移d) done = 0; tmp2 = []; while ~done tmp = [2*(r + w/2)*(rand(ts1,1)-0.5) -(r + w/2)*rand(ts1,1) + d]; tmp(:,3) = sqrt(tmp(:,1).^2 + (tmp(:,2)-d).^2); idx = find(tmp(:,3) > r - w/2 & tmp(:,3) < r + w/2); tmp2 = [tmp2; tmp(idx,1:2)]; if length(idx) >= ts done = 1; end end moon2 = tmp2(1:ts,:); % 合并数据和标签 data = [moon1; moon2]; labels = [zeros(ts,1); ones(ts,1)]; % 可视化双月数据 figure; scatter(moon1(:,1), moon1(:,2), 10, 'b', 'filled'); hold on; scatter(moon2(:,1), moon2(:,2), 10, 'r', 'filled'); xlabel('X1'); ylabel('X2'); title('双月数据集'); legend('类别0', '类别1'); grid on; end
运行这个函数后,你会得到包含1000个样本的两类双月数据,以及对应的标签。
第二步:直线分类的实现方法
需要先说明:原始双月数据集是线性不可分的——没有一条直线能完全分开所有样本,但我们可以找到一条直线尽可能减少分类错误,下面提供两种实用方法:
方法1:手动构造分割直线
观察双月的分布,两个月亮在y轴方向有偏移量d,我们可以直接取中间线作为分割直线,比如当d=-4时,直线y=-2就能很好地分开大部分样本。
实现代码:
% 生成双月数据 [data, labels] = dm(10,6,1000,-4); % 定义分割直线:y = d/2(这里d=-4,所以y=-2) x_range = linspace(min(data(:,1)), max(data(:,1)), 100); y_line = ones(size(x_range)) * (-4/2); % 可视化分类效果 figure; scatter(data(labels==0,1), data(labels==0,2),10,'b','filled'); hold on; scatter(data(labels==1,1), data(labels==1,2),10,'r','filled'); plot(x_range, y_line, 'k--', 'LineWidth',2); xlabel('X1'); ylabel('X2'); title('手动分割直线的双月分类'); legend('类别0','类别1','分割直线'); grid on; % 计算分类准确率 pred_labels = data(:,2) > (-4/2); % 直线上方为类别0,下方为类别1 accuracy = sum(pred_labels == labels)/length(labels); fprintf('手动分割直线的分类准确率:%.2f%%\n', accuracy*100);
方法2:用感知机算法训练线性分类器
感知机是一种简单的线性分类器,它会通过迭代优化找到一条近似最优的分割直线(因为数据线性不可分,我们设置最大迭代次数来终止训练)。
首先实现感知机函数:
function [w, b] = perceptron(data, labels, max_iter) if nargin <3, max_iter = 1000; end [n_samples, n_features] = size(data); w = zeros(n_features,1); b = 0; for iter = 1:max_iter misclassified = 0; for i = 1:n_samples x = data(i,:)'; y = labels(i); % 感知机判别规则:分类错误时更新参数 if y*(w'*x + b) <= 0 w = w + y*x; b = b + y; misclassified = misclassified +1; end end % 如果没有误分类样本,提前终止 if misclassified ==0 break; end end end
然后用感知机训练并可视化:
% 生成双月数据 [data, labels] = dm(10,6,1000,-4); % 把标签转换为-1/1(感知机的标准输入格式) labels(labels==0) = -1; % 训练感知机 [w, b] = perceptron(data, labels); % 计算分割直线的y值 x_range = linspace(min(data(:,1)), max(data(:,1)), 100); y_line = (-w(1)*x_range - b)/w(2); % 可视化 figure; scatter(data(labels==-1,1), data(labels==-1,2),10,'b','filled'); hold on; scatter(data(labels==1,1), data(labels==1,2),10,'r','filled'); plot(x_range, y_line, 'k--', 'LineWidth',2); xlabel('X1'); ylabel('X2'); title('感知机训练的分割直线'); legend('类别0','类别1','分割直线'); grid on; % 计算分类准确率 pred_scores = data*w + b; pred_labels = sign(pred_scores); pred_labels(pred_labels==-1) =0; labels(labels==-1)=0; accuracy = sum(pred_labels == labels)/length(labels); fprintf('感知机分类准确率:%.2f%%\n', accuracy*100);
补充说明
因为双月数据集天生线性不可分,所以以上两种方法都无法达到100%的分类准确率。如果想要完全正确分类,你需要使用非线性分类器(比如带核函数的SVM),但如果只要求用直线完成分类,上面的方法就足够满足需求了。
内容的提问来源于stack exchange,提问作者razin kac
相关产品推荐
相关产品推荐

