图像稀疏表示与字典训练:OMP + K-SVD(MATLAB 实现)
实现基于 OMP(正交匹配追踪) 的稀疏编码和 K-SVD(字典学习) 的字典训练,用于图像稀疏表示。
一、算法原理
1.1 稀疏表示模型
给定图像块集合 \(X \in \mathbb{R}^{n \times N}\)(\(n\) 为块维度,\(N\) 为块数量),稀疏表示的目标是找到字典 \(D \in \mathbb{R}^{n \times K}\) 和稀疏系数 \(\alpha \in \mathbb{R}^{K \times N}\),使得:
\[\min_{D,\alpha} \|X - D\alpha\|_F^2 \quad \text{s.t.} \quad \|\alpha_i\|_0 \leq T
\]
其中 \(T\) 为稀疏度,\(\|\cdot\|_0\) 为非零元素个数。
1.2 OMP 稀疏编码
OMP 是一种贪婪算法,通过迭代选择与残差最相关的原子,逐步构建稀疏解。
1.3 K-SVD 字典学习
K-SVD 通过交替优化:
- 稀疏编码阶段:固定字典 \(D\),用 OMP 求解 \(\alpha\)
- 字典更新阶段:固定 \(\alpha\),逐列更新字典原子 \(d_k\)
二、MATLAB 代码
2.1 主程序 sparse_representation_main.m
%% 图像稀疏表示与字典训练(OMP + K-SVD)
clear; clc; close all;
%% ========== 1. 加载图像并提取图像块 ==========
img = imread('cameraman.tif'); % 使用内置图像
if size(img,3) == 3
img = rgb2gray(img);
end
img = im2double(img);
% 提取图像块
patch_size = 8; % 8x8 图像块
overlap = 4; % 重叠像素
[X, positions] = extract_patches(img, patch_size, overlap);
[n, N] = size(X); % n=64, N=块数量
fprintf('提取 %d 个 %dx%d 图像块\n', N, patch_size, patch_size);
%% ========== 2. 初始化字典 ==========
K = 256; % 字典原子数量
D_init = initialize_dictionary(X, K);
fprintf('初始化字典: %d 个原子\n', K);
%% ========== 3. OMP 稀疏编码(使用初始字典) ==========
T = 10; % 稀疏度(每个块最多10个非零系数)
fprintf('执行 OMP 稀疏编码(稀疏度 T=%d)...\n', T);
Alpha_omp = zeros(K, N);
for i = 1:N
Alpha_omp(:,i) = omp(D_init, X(:,i), T);
end
% 重建误差
X_recon_omp = D_init * Alpha_omp;
rmse_omp = sqrt(mean((X(:) - X_recon_omp(:)).^2));
fprintf('OMP 重建 RMSE: %.4f\n', rmse_omp);
%% ========== 4. K-SVD 字典学习 ==========
max_iter = 50; % 最大迭代次数
fprintf('开始 K-SVD 字典学习(最大迭代 %d 次)...\n', max_iter);
[D_ksvd, Alpha_ksvd] = ksvd(X, D_init, T, max_iter);
% 重建误差
X_recon_ksvd = D_ksvd * Alpha_ksvd;
rmse_ksvd = sqrt(mean((X(:) - X_recon_ksvd(:)).^2));
fprintf('K-SVD 重建 RMSE: %.4f\n', rmse_ksvd);
%% ========== 5. 图像重建 ==========
fprintf('重建完整图像...\n');
% 使用 OMP 结果重建
img_recon_omp = reconstruct_image(Alpha_omp, D_init, positions, patch_size, size(img));
% 使用 K-SVD 结果重建
img_recon_ksvd = reconstruct_image(Alpha_ksvd, D_ksvd, positions, patch_size, size(img));
%% ========== 6. 结果可视化 ==========
figure('Position', [100, 100, 1400, 600]);
% 原始图像
subplot(2,3,1); imshow(img); title('原始图像');
% 字典可视化
subplot(2,3,2); visualize_dictionary(D_init, patch_size); title('初始字典(随机)');
subplot(2,3,3); visualize_dictionary(D_ksvd, patch_size); title('学习字典(K-SVD)');
% 重建图像
subplot(2,3,4); imshow(img_recon_omp);
title(sprintf('OMP 重建 (RMSE=%.4f)', rmse_omp));
subplot(2,3,5); imshow(img_recon_ksvd);
title(sprintf('K-SVD 重建 (RMSE=%.4f)', rmse_ksvd));
% 重建误差
subplot(2,3,6);
error_omp = abs(img - img_recon_omp);
error_ksvd = abs(img - img_recon_ksvd);
imshow(cat(3, error_omp, error_ksvd, zeros(size(error_omp))));
title('重建误差(红:OMP,绿:K-SVD)');
sgtitle('图像稀疏表示与字典训练(OMP + K-SVD)', 'FontSize', 14, 'FontWeight', 'bold');
%% ========== 7. 保存结果 ==========
save('sparse_representation_results.mat', 'D_init', 'D_ksvd', 'Alpha_omp', 'Alpha_ksvd', 'rmse_omp', 'rmse_ksvd');
fprintf('结果已保存到 sparse_representation_results.mat\n');
2.2 图像块提取函数 extract_patches.m
function [patches, positions] = extract_patches(img, patch_size, overlap)
% 从图像中提取重叠的图像块
[h, w] = size(img);
step = patch_size - overlap; % 步长
% 计算块数量
num_patches_h = floor((h - patch_size) / step) + 1;
num_patches_w = floor((w - patch_size) / step) + 1;
N = num_patches_h * num_patches_w;
patches = zeros(patch_size^2, N);
positions = zeros(N, 2);
idx = 1;
for i = 1:num_patches_h
for j = 1:num_patches_w
row_start = (i-1)*step + 1;
col_start = (j-1)*step + 1;
row_end = row_start + patch_size - 1;
col_end = col_start + patch_size - 1;
patch = img(row_start:row_end, col_start:col_end);
patches(:,idx) = patch(:);
positions(idx,:) = [row_start, col_start];
idx = idx + 1;
end
end
end
2.3 字典初始化函数 initialize_dictionary.m
function D = initialize_dictionary(X, K)
% 初始化字典:从数据随机选取或使用 DCT 字典
method = 'random'; % 'random' 或 'dct'
[n, N] = size(X);
if strcmp(method, 'random')
% 随机选择 K 个数据块作为初始字典
idx = randperm(N, K);
D = X(:,idx);
elseif strcmp(method, 'dct')
% 使用 DCT 字典
D = zeros(n, K);
for i = 1:K
% 生成 DCT 基向量
d = zeros(n,1);
for j = 1:n
d(j) = cos(pi*(i-1)*(j-1)/n);
end
D(:,i) = d / norm(d);
end
end
% 归一化字典原子
for i = 1:K
D(:,i) = D(:,i) / norm(D(:,i));
end
end
2.4 OMP 算法实现 omp.m
function alpha = omp(D, x, T)
% 正交匹配追踪(OMP)算法
% D: 字典 (n×K)
% x: 信号 (n×1)
% T: 稀疏度
% alpha: 稀疏系数 (K×1)
[n, K] = size(D);
alpha = zeros(K,1);
residual = x;
support = []; % 支持集
for t = 1:T
% 计算相关性
correlations = D' * residual;
% 选择最相关的原子
[~, idx] = max(abs(correlations));
support = union(support, idx);
% 最小二乘求解
D_support = D(:, support);
alpha_support = D_support \ x;
% 更新残差
residual = x - D_support * alpha_support;
% 检查残差范数
if norm(residual) < 1e-6
break;
end
end
% 填充稀疏系数
alpha(support) = alpha_support;
end
2.5 K-SVD 字典学习算法 ksvd.m
function [D, Alpha] = ksvd(X, D_init, T, max_iter)
% K-SVD 字典学习算法
% X: 数据矩阵 (n×N)
% D_init: 初始字典 (n×K)
% T: 稀疏度
% max_iter: 最大迭代次数
[n, N] = size(X);
K = size(D_init, 2);
D = D_init;
Alpha = zeros(K, N);
% 迭代
for iter = 1:max_iter
fprintf('K-SVD 迭代 %d/%d\n', iter, max_iter);
%% 1. 稀疏编码阶段
parfor i = 1:N
Alpha(:,i) = omp(D, X(:,i), T);
end
%% 2. 字典更新阶段
for k = 1:K
% 找到使用原子 k 的所有信号
usage = Alpha(k,:) ~= 0;
if sum(usage) == 0
continue; % 原子未被使用,跳过
end
% 计算误差矩阵
E = X(:,usage) - D * Alpha(:,usage);
E_k = E + D(:,k) * Alpha(k,usage); % 添加原子 k 的贡献
% 对误差矩阵进行 SVD
[U, S, V] = svd(E_k, 'econ');
% 更新原子 k 为第一个左奇异向量
D(:,k) = U(:,1);
% 更新对应的稀疏系数
Alpha(k,usage) = S(1,1) * V(:,1)';
end
% 归一化字典原子
for k = 1:K
D(:,k) = D(:,k) / norm(D(:,k));
end
% 计算当前重建误差
X_recon = D * Alpha;
rmse = sqrt(mean((X(:) - X_recon(:)).^2));
fprintf(' 迭代 %d: RMSE = %.4f\n', iter, rmse);
end
end
2.6 图像重建函数 reconstruct_image.m
function img_recon = reconstruct_image(Alpha, D, positions, patch_size, img_size)
% 使用稀疏系数和字典重建图像
[h, w] = img_size(1), img_size(2);
img_recon = zeros(h, w);
weight_map = zeros(h, w);
N = size(Alpha, 2);
idx = 1;
for i = 1:size(positions,1)
row_start = positions(i,1);
col_start = positions(i,2);
row_end = row_start + patch_size - 1;
col_end = col_start + patch_size - 1;
% 重建图像块
patch = reshape(D * Alpha(:,idx), patch_size, patch_size);
% 累加
img_recon(row_start:row_end, col_start:col_end) = ...
img_recon(row_start:row_end, col_start:col_end) + patch;
weight_map(row_start:row_end, col_start:col_end) = ...
weight_map(row_start:row_end, col_start:col_end) + 1;
idx = idx + 1;
end
% 加权平均
weight_map(weight_map == 0) = 1; % 避免除零
img_recon = img_recon ./ weight_map;
% 裁剪到 [0,1]
img_recon = max(0, min(1, img_recon));
end
2.7 字典可视化函数 visualize_dictionary.m
function visualize_dictionary(D, patch_size)
% 可视化字典原子
K = size(D, 2);
grid_size = ceil(sqrt(K));
figure;
for i = 1:K
subplot(grid_size, grid_size, i);
atom = reshape(D(:,i), patch_size, patch_size);
imshow(atom, []);
axis off;
end
end
三、运行说明
3.1 直接运行
- 将代码保存为
.m文件 - 确保 MATLAB 路径中有
cameraman.tif(或替换为你的图像) - 运行
sparse_representation_main.m
3.2 参数调优
| 参数 | 作用 | 建议值 |
|---|---|---|
patch_size |
图像块大小 | 8×8 或 16×16 |
K |
字典原子数量 | 256~1024 |
T |
稀疏度 | 5~20 |
max_iter |
K-SVD 迭代次数 | 20~50 |
overlap |
块重叠 | 4~6 |
3.3 预期结果
- OMP 重建:RMSE 约 0.05~0.1
- K-SVD 重建:RMSE 约 0.02~0.05
- 字典可视化:学习到的字典会显示边缘、角点等结构特征
四、算法优化建议
4.1 加速 OMP
% 使用 Cholesky 分解加速 OMP
function alpha = omp_cholesky(D, x, T)
R = D' * D;
residual = x;
support = [];
for t = 1:T
correlations = D' * residual;
[~, idx] = max(abs(correlations));
support = [support; idx];
L = chol(R(support,support), 'lower');
alpha_support = L' \ (L \ (D(:,support)' * x));
residual = x - D(:,support) * alpha_support;
end
end
4.2 改进 K-SVD
% 使用近似 K-SVD(MOD)加速
function D = approx_ksvd(X, D_init, T, max_iter)
for iter = 1:max_iter
% 稀疏编码
Alpha = zeros(size(D_init,2), size(X,2));
parfor i = 1:size(X,2)
Alpha(:,i) = omp(D_init, X(:,i), T);
end
% 字典更新(使用最小二乘)
D_init = X * Alpha' / (Alpha * Alpha' + 1e-6 * eye(size(Alpha,1)));
% 归一化
for k = 1:size(D_init,2)
D_init(:,k) = D_init(:,k) / norm(D_init(:,k));
end
end
D = D_init;
end
参考代码 图像的稀疏表示及字典训练代码OMP和KSVD www.youwenfan.com/contentcnw/81798.html
五、应用场景
| 应用 | 说明 |
|---|---|
| 图像去噪 | 稀疏表示能有效分离信号和噪声 |
| 图像修复 | 缺失像素可通过稀疏编码恢复 |
| 压缩感知 | 稀疏表示是压缩感知的基础 |
| 特征提取 | 字典原子可作为图像特征 |
浙公网安备 33010602011771号