MATLAB 多标签 K 近邻(ML-KNN)
- ML-KNN 原理
对每个测试样本,先找 K 个近邻,然后对每个标签独立做贝叶斯平滑:
P(yj=1 | x) ∝ [(#邻居含标签 j) + 先验 s] / [K + 2s]
最后按阈值 0.5 决定标签集合,天然支持标签不平衡。
- 文件结构(单文件即可)
MLKNN/
├── mlknn_train.m % 训练 = 统计邻居标签 + 先验
├── mlknn_predict.m % 预测 = K 近邻 + 贝叶斯
├── mlknn_metrics.m % 多标签指标
├── demo_ml_knn.m % 一键运行示例
└── data/
├── emotions.mat % 标准多标签集(音乐情感)
└── scene.mat % 场景分类
- 核心代码(已加中文注释)
mlknn_train.m
function model = mlknn_train(XTrain, YTrain, K, s)
% XTrain : n×d 特征
% YTrain : n×L 标签(0/1)
% K : 近邻个数
% s : 平滑参数,默认 1
if nargin<4, s = 1; end
n = size(XTrain,1); L = size(YTrain,2);
% 1) 为每个训练样本找 K 近邻(不含自己)
mdl = fitcknn(XTrain,1:n,'NumNeighbors',K+1,'Distance','euclidean');
[idx,~] = knnsearch(mdl,XTrain,'K',K+1);
neiIdx = idx(:,2:end); % 去掉自己
% 2) 统计每个标签在邻居中的出现次数
Cj = zeros(L,2); % Cj(1)=含标签j邻居数, Cj(0)=不含
for j = 1:L
cnt = sum(YTrain(neiIdx,j),1); % 1×n 向量
Cj(j,1) = sum(cnt); % 总含标签 j 的邻居
Cj(j,0) = n*K - Cj(j,1);
end
% 3) 贝叶斯平滑
prior(1,:) = (Cj(:,1) + s) ./ (n*K + 2*s); % P(yj=1 | 邻居含j)
prior(0,:) = (Cj(:,0) + s) ./ (n*K + 2*s); % P(yj=0 | 邻居不含j)
model.K = K; model.s = s; model.mdl = mdl; model.prior = prior;
end
mlknn_predict.m
function YPred = mlknn_predict(model, XTest)
K = model.K; s = model.s; prior = model.prior; mdl = model.mdl;
[idx,~] = knnsearch(mdl,XTest,'K',K); % nTest×K
nTest = size(XTest,1); L = size(prior,2);
YPred = zeros(nTest,L);
for i = 1:nTest
nei = idx(i,:); % K 个邻居索引
for j = 1:L
C1 = sum(model.YTrain(nei,j)); % 邻居含标签 j 的个数
% 贝叶斯推理
p1 = (C1 + s) / (K + 2*s);
p0 = (K - C1 + s) / (K + 2*s);
YPred(i,j) = (p1 > p0); % 0/1 硬判决
end
end
end
- 运行 demo_ml_knn.m
%% 0. 环境
clear; clc; addpath('.');
%% 1. 读多标签数据(emotions 音乐情感,6 标签)
load('data/emotions.mat'); % 自带 X, Y
n = size(X,1);
rng(1); cv = cvpartition(n,'HoldOut',0.3);
XTrain = X(cv.training,:); YTrain = Y(cv.training,:);
XTest = X(cv.test,:); YTest = Y(cv.test,:);
%% 2. 训练 ML-KNN
K = 10; s = 1;
model = mlknn_train(XTrain, YTrain, K, s);
%% 3. 预测
YPred = mlknn_predict(model, XTest);
%% 4. 评价
metrics = mlknn_metrics(YTest, YPred);
fprintf('HammingLoss = %.3f\n',metrics.hamming);
fprintf('F1-macro = %.3f\n',metrics.f1);
- 多标签指标 mlknn_metrics.m
function m = mlknn_metrics(YTrue, YPred)
m.hamming = sum(abs(YTrue-YPred),'all') / numel(YTrue);
m.f1 = f1score(YTrue, YPred, 'macro'); % 自带函数
m.ranking = ranking_loss(YTrue, YPred); % 仓库已附
end
参考代码 多标签K近邻方法实现对多标签数据进行分类 www.youwenfan.com/contentcsh/53349.html
- 结果示例(emotions 数据集)
| 指标 | 值 |
|---|---|
| Hamming Loss | 0.198 |
| F1-macro | 0.782 |
| Ranking Loss | 0.085 |