Add basic code

This commit is contained in:
Ksyer
2024-07-16 15:55:00 +08:00
parent e98983b374
commit 1c009370a5
6 changed files with 121 additions and 19 deletions
Executable
+27
View File
@@ -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
+6 -5
View File
@@ -38,7 +38,6 @@ function [x_hat_wl, x_hat_d] = cVAMPro(y, A, lambda, tau, Kit)
h_1 = h_1_next; h_1 = h_1_next;
Q_1 = Q_1_next; Q_1 = Q_1_next;
end end
end end
% SoftThreshold function % SoftThreshold function
@@ -48,7 +47,8 @@ function x = ST(h_1, lambda, Q_1)
for i = 1:N for i = 1:N
sign = h_1(i) ./ abs(h_1(i)); 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); x(i) = sign .* (diff ./ Q_1) .* SF(diff);
end end
@@ -73,9 +73,10 @@ function chi_1 = F1(x_1, lambda, Q_1)
count = 0; count = 0;
for i = 1:N for i = 1:N
temp = Q_1 * abs(x_1(i)) + lambda(i); % 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) * 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;
count = count + (2 - lambda / temp) * SF(abs(x_1(i)));
end end
chi_1 = count / (2 * N * Q_1); chi_1 = count / (2 * N * Q_1);
+21 -13
View File
@@ -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); sz = size(A);
N = sz(2); N = sz(2);
LASSO_lambda = 2;
gamma = sz(1) / sz(2); gamma = sz(1) / sz(2);
@@ -10,22 +14,26 @@ function [x_hat, threshold] = debiased_LASSO(A, y_noise, P_fa, noise_sigma_2)
minimize(LASSO_lambda * norm(x_LASSO, 1) + norm(y_noise - A * x_LASSO, 2)) minimize(LASSO_lambda * norm(x_LASSO, 1) + norm(y_noise - A * x_LASSO, 2))
cvx_end cvx_end
rho_active = sum(abs(x_LASSO) > 1e-3)/N; rho_active = sum(abs(x_LASSO) > 1e-3) / N;
Q_hat = (gamma - rho_active)/(1 - rho_active); 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 = sum((abs(x_LASSO) > 1e-3) .* (2 - LASSO_lambda ./ (Q_hat * abs(x_LASSO) + LASSO_lambda))) / 2 / N;
diff = 1; diff = 1;
while(diff > 1e-4)
while (diff > 1e-4)
Rho_pre = Rho; 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); diff = abs(Rho - Rho_pre);
end 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; x_hat = x_d_CROD;
RSS = 1/sz(2) * norm(y_noise - A*x_LASSO, 2) ^ 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; sigma_w_2 = (gamma * (1 - gamma)) / ((gamma - Rho) ^ 2) * RSS + noise_sigma_2;
threshold = -sigma_w_2 * log(P_fa); 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 end
+35
View File
@@ -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
+11
View File
@@ -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
+20
View File
@@ -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