From 1c009370a5877765d78a62218b9e0cd47dec6e45 Mon Sep 17 00:00:00 2001 From: Ksyer <> Date: Tue, 16 Jul 2024 15:55:00 +0800 Subject: [PATCH] Add basic code --- basic/FISTA.m | 27 +++++++++++++++++++++++++++ basic/cVAMPro.m | 11 ++++++----- basic/debiased_LASSO.m | 36 ++++++++++++++++++++++-------------- basic/debiased_LASSO_FISTA.m | 35 +++++++++++++++++++++++++++++++++++ basic/get_noise.m | 11 +++++++++++ basic/sft_thd.m | 20 ++++++++++++++++++++ 6 files changed, 121 insertions(+), 19 deletions(-) create mode 100755 basic/FISTA.m create mode 100755 basic/debiased_LASSO_FISTA.m create mode 100755 basic/get_noise.m create mode 100755 basic/sft_thd.m diff --git a/basic/FISTA.m b/basic/FISTA.m new file mode 100755 index 0000000..24f7223 --- /dev/null +++ b/basic/FISTA.m @@ -0,0 +1,27 @@ +function [z] = FISTA(y, A, lambda, delta) + +x_pre = A'*y; +t = 1; +z = x_pre; +z_pre = z; +t_pre = t; +N = size(A, 2); +diff = 1; +E = eig(A'*A); +L = E(end); +temp1 = A'*y/L; +temp2 = eye(N) - A'*A/L; +k = 0; +while((diff > delta) && (k < 1000)) + temp = temp1 + temp2 * z_pre; + x = sft_thd(temp, lambda/L); + t = 0.5*(1 + sqrt(1+4*t_pre*t_pre)); + z = x + (x - x_pre) * (t_pre-1) / t; + diff = mean(abs(z_pre - z)); + x_pre = x; + z_pre = z; + t_pre = t; + k = k + 1; +end + +end \ No newline at end of file diff --git a/basic/cVAMPro.m b/basic/cVAMPro.m index ee21f8d..3d8a420 100755 --- a/basic/cVAMPro.m +++ b/basic/cVAMPro.m @@ -38,7 +38,6 @@ function [x_hat_wl, x_hat_d] = cVAMPro(y, A, lambda, tau, Kit) h_1 = h_1_next; Q_1 = Q_1_next; end - end % SoftThreshold function @@ -48,7 +47,8 @@ function x = ST(h_1, lambda, Q_1) for i = 1:N sign = h_1(i) ./ abs(h_1(i)); - diff = abs(h_1(i)) - lambda(i); + % diff = abs(h_1(i)) - lambda(i); + diff = abs(h_1(i)) - lambda; x(i) = sign .* (diff ./ Q_1) .* SF(diff); end @@ -73,9 +73,10 @@ function chi_1 = F1(x_1, lambda, Q_1) count = 0; for i = 1:N - temp = Q_1 * abs(x_1(i)) + lambda(i); - count = count + (2 - lambda(i) / temp) * SF(abs(x_1(i))); - % count = count + (2-lambda(i)/temp) * (abs(x_1(i)) > 1e-4); + % temp = Q_1 * abs(x_1(i)) + lambda(i); + % count = count + (2 - lambda(i) / temp) * SF(abs(x_1(i))); + temp = Q_1 * abs(x_1(i)) + lambda; + count = count + (2 - lambda / temp) * SF(abs(x_1(i))); end chi_1 = count / (2 * N * Q_1); diff --git a/basic/debiased_LASSO.m b/basic/debiased_LASSO.m index 6f73661..601f9dd 100755 --- a/basic/debiased_LASSO.m +++ b/basic/debiased_LASSO.m @@ -1,7 +1,11 @@ -function [x_hat, threshold] = debiased_LASSO(A, y_noise, P_fa, noise_sigma_2) +function [x_hat, sigma_w_2, threshold] = debiased_LASSO(A, y_noise, P_fa, noise_sigma_2, LASSO_lambda) + + if nargin < 5 + LASSO_lambda = 1e-1; + end + sz = size(A); N = sz(2); - LASSO_lambda = 2; gamma = sz(1) / sz(2); @@ -9,23 +13,27 @@ function [x_hat, threshold] = debiased_LASSO(A, y_noise, P_fa, noise_sigma_2) variable x_LASSO(N) complex minimize(LASSO_lambda * norm(x_LASSO, 1) + norm(y_noise - A * x_LASSO, 2)) cvx_end - - rho_active = sum(abs(x_LASSO) > 1e-3)/N; - Q_hat = (gamma - rho_active)/(1 - rho_active); - Rho = sum((abs(x_LASSO) > 1e-3).* (2 - LASSO_lambda./(Q_hat*abs(x_LASSO) + LASSO_lambda))) / 2 / N; + + rho_active = sum(abs(x_LASSO) > 1e-3) / N; + Q_hat = (gamma - rho_active) / (1 - rho_active); + Rho = sum((abs(x_LASSO) > 1e-3) .* (2 - LASSO_lambda ./ (Q_hat * abs(x_LASSO) + LASSO_lambda))) / 2 / N; diff = 1; - while(diff > 1e-4) + + while (diff > 1e-4) Rho_pre = Rho; - Rho = sum((abs(x_LASSO) > 1e-3).* (2 - LASSO_lambda./((gamma-Rho)/(1-Rho)*abs(x_LASSO) + LASSO_lambda))) / 2 / N; + Rho = sum((abs(x_LASSO) > 1e-3) .* (2 - LASSO_lambda ./ ((gamma - Rho) / (1 - Rho) * abs(x_LASSO) + LASSO_lambda))) / 2 / N; diff = abs(Rho - Rho_pre); end - Q_hat = (gamma-Rho)/(1-Rho); - x_d_CROD = x_LASSO + A'*(y_noise - A*x_LASSO)/Q_hat; + + Q_hat = (gamma - Rho) / (1 - Rho); + x_d_CROD = x_LASSO + A' * (y_noise - A * x_LASSO) / Q_hat; x_hat = x_d_CROD; - RSS = 1/sz(2) * norm(y_noise - A*x_LASSO, 2) ^ 2; - sigma_w_2 = (gamma * (1 - gamma)) / ((gamma - Rho)^2) * RSS + noise_sigma_2; + RSS = 1 / sz(1) * norm(y_noise - A * x_LASSO, 2) ^ 2; + sigma_w_2 = (gamma * (1 - gamma)) / ((gamma - Rho) ^ 2) * RSS + noise_sigma_2; threshold = -sigma_w_2 * log(P_fa); - + xx = x_LASSO; + FAR_M = 1 / gamma; + xx(FAR_M+1: 2*FAR_M) = xx(FAR_M+1: 2*FAR_M) - 1; + xx = real(xx); end - diff --git a/basic/debiased_LASSO_FISTA.m b/basic/debiased_LASSO_FISTA.m new file mode 100755 index 0000000..80b79b1 --- /dev/null +++ b/basic/debiased_LASSO_FISTA.m @@ -0,0 +1,35 @@ +function [x_hat, sigma_w_2, threshold] = debiased_LASSO_FISTA(A, y_noise, P_fa, noise_sigma_2, LASSO_lambda) + + if nargin < 5 + LASSO_lambda = 1e-1; + end + + sz = size(A); + N = sz(2); + + gamma = sz(1) / sz(2); + + x_LASSO = FISTA(y_noise, A, LASSO_lambda, 1e-5); + + rho_active = sum(abs(x_LASSO) > 1e-3) / N; + Q_hat = (gamma - rho_active) / (1 - rho_active); + Rho = sum((abs(x_LASSO) > 1e-3) .* (2 - LASSO_lambda ./ (Q_hat * abs(x_LASSO) + LASSO_lambda))) / 2 / N; + diff = 1; + + while (diff > 1e-4) + Rho_pre = Rho; + Rho = sum((abs(x_LASSO) > 1e-3) .* (2 - LASSO_lambda ./ ((gamma - Rho) / (1 - Rho) * abs(x_LASSO) + LASSO_lambda))) / 2 / N; + diff = abs(Rho - Rho_pre); + end + + Q_hat = (gamma - Rho) / (1 - Rho); + x_d_CROD = x_LASSO + A' * (y_noise - A * x_LASSO) / Q_hat; + x_hat = x_d_CROD; + + RSS = 1 / sz(1) * norm(y_noise - A * x_LASSO, 2) ^ 2; + sigma_w_2 = (gamma * (1 - gamma)) / ((gamma - Rho) ^ 2) * RSS + noise_sigma_2; + threshold = -sigma_w_2 * log(P_fa); + xx = x_LASSO; + FAR_M = 1 / gamma; + mean_xx = mean(xx(FAR_M+1: 2*FAR_M)); +end diff --git a/basic/get_noise.m b/basic/get_noise.m new file mode 100755 index 0000000..191b4d1 --- /dev/null +++ b/basic/get_noise.m @@ -0,0 +1,11 @@ +function noise = get_noise(noise_sigma, N, M, type) + if nargin < 4 + type = "complex"; + end + + if type == "complex" + noise = random('Normal', 0, noise_sigma / sqrt(2), N, M) + 1j * random('Normal', 0, noise_sigma / sqrt(2), N, M); + else + noise = random('Normal', 0, noise_sigma, N, M); + end +end \ No newline at end of file diff --git a/basic/sft_thd.m b/basic/sft_thd.m new file mode 100755 index 0000000..95e46ba --- /dev/null +++ b/basic/sft_thd.m @@ -0,0 +1,20 @@ +function y = sft_thd(x, thd) + +if isequal(size(x), size(thd)) + + tmp = abs(x); + y = x; + y(tmp <= thd) = 0; + y(tmp > thd) = (tmp(tmp > thd) - thd(tmp > thd)) .* x(tmp > thd) ./ tmp(tmp > thd); + +else + + tmp = abs(x); + y = x; + y(tmp <= thd) = 0; + y(tmp > thd) = (tmp(tmp > thd) - thd) .* x(tmp > thd) ./ tmp(tmp > thd); + +end + + +end \ No newline at end of file