基于MATLAB的SMOTE算法实现
一、SMOTE算法原理
核心思想:通过插值生成少数类新样本,解决数据不平衡问题
关键步骤:
- 确定少数类样本:识别数据集中样本量最少的类别
- 计算k近邻:对每个少数类样本计算k个最近邻
- 生成新样本:在样本与近邻间随机插值
- 数据集重构:合并原始数据与合成样本
二、MATLAB实现代码
2.1 基础二分类实现
function [X_resampled, y_resampled] = mySMOTE(X, y, k, ratio)
% 参数说明:
% X: 特征矩阵 (n_samples × n_features)
% y: 标签向量 (n_samples × 1)
% k: 近邻数
% ratio: 过采样倍数
% 找出少数类样本
minority_class = mode(y);
idx = y == minority_class;
X_minority = X(idx, :);
n_minority = size(X_minority, 1);
% 计算k近邻
[idx_knn, ~] = knnsearch(X_minority, X_minority, 'K', k+1);
neighbors = idx_knn(:, 2:end); % 排除自身
% 生成新样本
n_synthetic = round(ratio * n_minority);
synthetic_samples = zeros(n_synthetic, size(X, 2));
for i = 1:n_synthetic
% 随机选择基样本
base_idx = randi(n_minority);
% 随机选择近邻
neighbor_idx = neighbors(base_idx, randi(k));
% 线性插值
diff = X_minority(neighbor_idx, :) - X_minority(base_idx, :);
gap = rand(1, size(X, 2));
synthetic_samples(i, :) = X_minority(base_idx, :) + gap .* diff;
end
% 合并数据集
X_resampled = [X; synthetic_samples];
y_resampled = [y; repmat(minority_class, n_synthetic, 1)];
end
2.2 多分类扩展实现
function [X_resampled, y_resampled] = multiClassSMOTE(X, y, k, ratio)
% 处理多分类不平衡问题
classes = unique(y);
X_resampled = X;
y_resampled = y;
for i = 1:length(classes)
class = classes(i);
idx = y == class;
n_class = sum(idx);
% 仅处理少数类
if n_class < mean(y == classes) && n_class > 1
X_minority = X(idx, :);
[X_synthetic, y_synthetic] = mySMOTE(X_minority, class, k, ratio);
X_resampled = [X_resampled; X_synthetic];
y_resampled = [y_resampled; y_synthetic];
end
end
end
三、参数优化
| 参数 | 推荐范围 | 影响分析 |
|---|---|---|
| k值 | 3-10 | 过小导致噪声敏感,过大降低多样性 |
| 过采样倍数 | 1-5 | 倍数过高导致过拟合 |
| 距离度量 | 欧氏/马氏 | 影响近邻选择准确性 |
| 平衡策略 | 过采样/欠采样 | 需结合欠采样避免类别反转 |
四、性能优化
4.1 数据预处理
% 特征标准化
scaler = fitpreprocess(X);
X_scaled = transform(scaler, X);
% 处理缺失值
X_clean = fillmissing(X, 'linear');
4.2 并行计算加速
% 启用并行池
parpool('local');
% 并行生成样本
parfor i = 1:num_cores
% 分块处理数据
end
4.3 边界样本保护
% Tomek Links检测
[idx_tomek] = detect_tomek_links(X, y, k);
X_protected = X(~idx_tomek, :);
五、应用案例
5.1 医疗诊断数据集
% 加载数据
load('medical_data.mat');
X = data(:, 1:end-1);
y = data(:, end);
% 执行SMOTE
[X_balanced, y_balanced] = multiClassSMOTE(X, y, 5, 2);
% 模型训练
model = fitcsvm(X_balanced, y_balanced);
5.2 工业故障检测
% 参数设置
k = 7; % 高维数据需增大k值
ratio = 3; % 严重不平衡时提高倍数
% 执行过采样
[X_resampled, y_resampled] = mySMOTE(X_fault, y_fault, k, ratio);
% 交叉验证
cv = cvpartition(y_resampled, 'KFold', 5);
accuracy = crossval(@(XTrain,yTrain,XTest,yTest) ...
sum(predict(fitcsvm(XTrain,yTrain), XTest) == yTest)/numel(yTest), ...
X_resampled, y_resampled, 'partition', cv);
disp(['平均准确率: ', num2str(mean(accuracy))]);
参考代码 matlab下的实现smote算法 www.youwenfan.com/contentcni/65444.html
六、可视化分析
6.1 二维数据分布
% 生成示例数据
[X, y] = make_classification(1000, 2, 1, 2, 1, 0.5);
% 执行SMOTE
[X_balanced, y_balanced] = mySMOTE(X, y, 5, 2);
% 绘制分布
figure;
gscatter(X(:,1), X(:,2), y);
hold on;
gscatter(X_balanced(:,1), X_balanced(:,2), y_balanced);
title('SMOTE效果可视化');
legend('原始少数类', '原始多数类', '合成少数类');
6.2 决策边界对比
% 训练模型
model_original = fitcsvm(X, y);
model_balanced = fitcsvm(X_balanced, y_balanced);
% 生成网格
[x1Grid, x2Grid] = meshgrid(linspace(-5,5,50));
xGrid = [x1Grid(:), x2Grid(:)];
% 预测概率
[~,score_original] = predict(model_original, xGrid);
[~,score_balanced] = predict(model_balanced, xGrid);
% 绘制决策面
figure;
contourf(x1Grid, x2Grid, reshape(score_original(:,2),50,50), 'LineColor', 'none');
hold on;
gscatter(X(:,1), X(:,2), y);
title('原始数据决策边界');
figure;
contourf(x1Grid, x2Grid, reshape(score_balanced(:,2),50,50), 'LineColor', 'none');
hold on;
gscatter(X_balanced(:,1), X_balanced(:,2), y_balanced);
title('SMOTE平衡后决策边界');