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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:39:17