Files
FAR_CS/wide vs narrow/single_point_test/test_lambda.m
T
2024-07-22 21:17:57 +08:00

103 lines
3.0 KiB
Matlab
Executable File

clc; clear; close all;
N = 256;
M = 16;
MN = N * M;
sigmas = [0.01];
% noise_sigma = 0.01;
filename = "./results/Signal_Model_" + string(N) + "_" + string(M) + ".mat";
curr_metric = N;
for TT = 1: 100
TT
tau = 1e-6;
iter_max = 25;
epi = 0;
[C_n, A] = get_Psi(N, M, epi);
betas_wide = zeros(M * N, 1) + 1;
[Lambda, Lambda_C, x] = get_sparse_vector(N, M, betas_wide, true);
lg_lambdas = -3: 0.05: -1;
lambdas = 10 .^ lg_lambdas;
method = "cVAMPro";
for sigma_idx = 1: length(sigmas)
sigma = sigmas(sigma_idx);
noise = get_noise(sigma, N, 1);
y_noise = A * x + noise;
MSEs = zeros(length(lambdas), 1);
all_x_hat = zeros(length(x), length(lambdas));
for lambda_idx = 1: length(lambdas)
lambda = lambdas(lambda_idx);
if method == "cVAMPro"
[x_LASSO, x_hat_d] = cVAMPro(y_noise, A, lambda, tau, iter_max);
elseif method == "debiased\_LASSO"
[x_LASSO, sigma_w_2, threshold] = debiased_LASSO(A, y_noise, 0, sigma^2, lambda);
elseif method == "LASSO"
cvx_begin quiet
variable x_LASSO(MN) complex
minimize(lambda * norm(x_LASSO, 1) + norm(y_noise - A * x_LASSO, 2))
cvx_end
elseif method == "FISTA"
x_LASSO = FISTA(y_noise, A, LASSO_lambda, 1e-5);
elseif method == "debiased\_LASSO\_FISTA"
[x_LASSO, sigma_w_2, threshold] = debiased_LASSO_FISTA(A, y_noise, 0, noise_sigma^2, lambda);
elseif method == "BP"
sz = size(A);
N = sz(2);
cvx_begin quiet
variable x_LASSO(N) complex
minimize(norm(x_LASSO, 1))
subject to
A * x_LASSO == y_noise
cvx_end
end
MSEs(lambda_idx) = sum(real(x_LASSO - x));
all_x_hat(:, lambda_idx) = x_LASSO;
end
abs_MSEs = abs(MSEs);
[metric, idx] = min(abs_MSEs);
if curr_metric > metric
ref_lambda = lambdas(idx);
save(filename, "A", "C_n", "ref_lambda");
curr_metric = metric;
end
if curr_metric < 0.1
figure;
subplot(211);
semilogx(lambdas, MSEs);
yline(0);
ylim([-length(Lambda) - 0.1, 0.1]);
xlabel("\lambda");
ylabel("MSE (x\_hat - x)");
title("N = " + string(N) + ", M = " + string(M) + ", \sigma = " + string(sigma) + ", " + method);
subplot(212);
plot(C_n);
xlabel("Frequency code");
ylabel("Times");
break
end
% for i = 1: length(C_n)
% fprintf("%d, ", C_n(i))
% end
% fprintf("%.15f\n", sum(MSEs));
% filename = ...
% string(N) + "_" + ...
% string(M) + "_" + ...
% string(sigma) + "_" + ...
% method + ".mat";
% save(filename, "lambdas", "MSEs", "x", "all_x_hat");
% figure;
% semilogx(lambdas, xx);
% yline(0);
% ylim([-length(Lambda) - 0.1, 0.1]);
% xlabel("\lambda");
% ylabel("MSE (x\_hat - x)");
% title("N = " + string(FAR_N) + ", M = " + string(FAR_M) + ", \sigma = " + string(sigma) + ", " + method);
end
end