Update wide vs narrow code
This commit is contained in:
@@ -0,0 +1,196 @@
|
||||
%% Initial
|
||||
clc; clear; close all;
|
||||
|
||||
trail_times = 250;
|
||||
method = "cVAMPro";
|
||||
Ns = [64, 128, 256, 512, 1024];
|
||||
Ms = [4, 8, 16, 32];
|
||||
sigma = 0.01;
|
||||
tau = 1e-6;
|
||||
iter_max = 100;
|
||||
epi = 0;
|
||||
|
||||
lg_lambdas = -5: 0.05: -1;
|
||||
lambdas = 10 .^ lg_lambdas;
|
||||
|
||||
H0_REE_means = zeros(length(Ns), length(Ms));
|
||||
H0_REE_maxs = zeros(length(Ns), length(Ms));
|
||||
H1_REE_means = zeros(length(Ns), length(Ms));
|
||||
H1_REE_maxs = zeros(length(Ns), length(Ms));
|
||||
|
||||
figure;
|
||||
|
||||
for N_idx = 1: length(Ns)
|
||||
for M_idx = 1: length(Ms)
|
||||
subplot(length(Ns), length(Ms), (N_idx - 1) * length(Ms) + M_idx);
|
||||
FAR_N = Ns(N_idx);
|
||||
FAR_M = Ms(M_idx);
|
||||
|
||||
fprintf("Simluating: N = %d, M = %d\n\n", FAR_N, FAR_M);
|
||||
|
||||
MN = FAR_N * FAR_M;
|
||||
betas_wide = zeros(FAR_M * FAR_N, 1) + 1;
|
||||
[Lambda, Lambda_C, x] = get_sparse_vector(FAR_N, FAR_M, betas_wide, true);
|
||||
|
||||
% get lambda
|
||||
signal_model_filename = "./Signal_Model/Signal_Model_" + string(FAR_N) + "_" + string(FAR_M) + ".mat";
|
||||
if exist(signal_model_filename, "file")
|
||||
load(signal_model_filename, "A", "C_n", "ref_lambda");
|
||||
else
|
||||
curr_metric = FAR_N;
|
||||
for TT = 1: 20
|
||||
[C_n, A] = get_Psi(FAR_N, FAR_M, epi);
|
||||
|
||||
noise = get_noise(sigma, FAR_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);
|
||||
[x_LASSO, x_hat_d] = cVAMPro(y_noise, A, lambda, tau, iter_max);
|
||||
MSEs(lambda_idx) = sum(real(x_LASSO - x));
|
||||
all_x_hat(:, lambda_idx) = x_LASSO;
|
||||
end
|
||||
|
||||
[metric, idx] = min(abs(MSEs));
|
||||
if curr_metric > metric
|
||||
ref_lambda = lambdas(idx);
|
||||
save(signal_model_filename, "A", "C_n", "ref_lambda");
|
||||
curr_metric = metric;
|
||||
end
|
||||
end
|
||||
end
|
||||
fprintf("get lambda: lambda = %.15f. \n", ref_lambda);
|
||||
fprintf("Signal Model filename: " + signal_model_filename + "\n\n");
|
||||
|
||||
% train
|
||||
LASSO_lambda = ref_lambda;
|
||||
simulate_results_filename = ...
|
||||
"./data5/" + ...
|
||||
"FAR" + "_" + ...
|
||||
string(FAR_N) + "_" + ...
|
||||
string(FAR_M)+ "_" + ...
|
||||
method + "_" + ...
|
||||
string(sigma) + "_" + ...
|
||||
trail_times + ...
|
||||
".mat";
|
||||
if exist(simulate_results_filename, "file")
|
||||
load( ...
|
||||
simulate_results_filename, ...
|
||||
"recovery_results", "sigma_w2", "thresholds",...
|
||||
"C_n", "A", "Lambda", "Lambda_C", "x" ...
|
||||
);
|
||||
else
|
||||
% Mento Carlo Recovery
|
||||
[recovery_results, sigma_w2, thresholds] = All_Recovery2(A, x, LASSO_lambda, sigma, trail_times, LASSO_lambda, tau);
|
||||
|
||||
save( ...
|
||||
simulate_results_filename, ...
|
||||
"recovery_results", "sigma_w2", "thresholds",...
|
||||
"C_n", "A", "Lambda", "Lambda_C", "x" ...
|
||||
);
|
||||
end
|
||||
fprintf("Simulate results filename: " + simulate_results_filename + "\n\n");
|
||||
|
||||
% test
|
||||
i1 = 1; i2 = 1;
|
||||
Lambda_distributes = zeros(2, length(Lambda) * trail_times);
|
||||
Lambda_C_distributes = zeros(2, length(Lambda_C) * trail_times);
|
||||
|
||||
for t = 1:trail_times
|
||||
x_0_hat = squeeze(recovery_results(1, t, :));
|
||||
x_1_hat = squeeze(recovery_results(2, t, :)) - x;
|
||||
if anynan(x_1_hat)
|
||||
continue;
|
||||
end
|
||||
|
||||
for i = 1: length(Lambda) + length(Lambda_C)
|
||||
if ismember(i, Lambda_C)
|
||||
Lambda_C_distributes(1, i1) = x_0_hat(i);
|
||||
Lambda_C_distributes(2, i1) = x_1_hat(i);
|
||||
i1 = i1 + 1;
|
||||
elseif ismember(i, Lambda)
|
||||
Lambda_distributes(1, i2) = x_0_hat(i);
|
||||
Lambda_distributes(2, i2) = x_1_hat(i);
|
||||
i2 = i2 + 1;
|
||||
end
|
||||
end
|
||||
end
|
||||
real_H00 = real(Lambda_C_distributes(1, 1:i1-1));
|
||||
real_H01 = real(Lambda_distributes(1, 1:i2-1));
|
||||
real_H10 = real(Lambda_C_distributes(2, 1:i1-1));
|
||||
real_H11 = real(Lambda_distributes(2, 1:i2-1));
|
||||
|
||||
% figure;
|
||||
% subplot(2, 2, 1); histfit(real_H00); xlabel("x\_hat"); ylabel("times"); title("H_0 (not in support set)")
|
||||
% subplot(2, 2, 2); histfit(real_H01); xlabel("x\_hat"); ylabel("times"); title("H_0 (in support set)")
|
||||
% subplot(2, 2, 3); histfit(real_H10); xlabel("x\_hat"); ylabel("times"); title("H_1 (not in support set)")
|
||||
% subplot(2, 2, 4); histfit(real_H11); xlabel("x\_hat"); ylabel("times"); title("H_1 (in support set)")
|
||||
% sgtitle("N = " + string(FAR_N) + ", M = " + string(FAR_M));
|
||||
|
||||
histfit(real_H11); xlabel("x\_hat"); ylabel("times"); title("N = " + string(FAR_N) + ", M = " + string(FAR_M));
|
||||
|
||||
means = [ ...
|
||||
mean(Lambda_C_distributes(1, 1:i1-1)), ...
|
||||
mean(Lambda_distributes(1, i2-1)), ...
|
||||
mean(Lambda_C_distributes(2, 1:i1-1)), ...
|
||||
mean(Lambda_distributes(2, i2-1)), ...
|
||||
mean([Lambda_distributes(1, i2-1) Lambda_C_distributes(1, 1:i1-1)]) ...
|
||||
];
|
||||
|
||||
stds = [ ...
|
||||
std(Lambda_C_distributes(1, 1:i1-1)), ...
|
||||
std(Lambda_distributes(1, i2-1)), ...
|
||||
std(Lambda_C_distributes(2, 1:i1-1)), ...
|
||||
std(Lambda_distributes(2, i2-1)), ...
|
||||
std([Lambda_distributes(1, i2-1) Lambda_C_distributes(1, 1:i1-1)]) ...
|
||||
];
|
||||
|
||||
|
||||
fprintf("H_00: mu = %.15f, std = %.15f\n", means(1), stds(1));
|
||||
fprintf("H_01: mu = %.15f, std = %.15f\n", means(2), stds(2));
|
||||
fprintf("H_10: mu = %.15f, std = %.15f\n", means(3), stds(3));
|
||||
fprintf("H_11: mu = %.15f, std = %.15f\n", means(4), stds(4));
|
||||
fprintf("H_0: mu = %.15f, std = %.15f\n", means(5), stds(5));
|
||||
|
||||
[ss, vs, H0_mean, H0_max, H1_mean, H1_max] = get_sigma_var(recovery_results, sigma_w2, x);
|
||||
H0_REE_means(N_idx, M_idx) = H0_mean;
|
||||
H0_REE_maxs(N_idx, M_idx) = H0_max;
|
||||
H1_REE_means(N_idx, M_idx) = H1_mean;
|
||||
H1_REE_maxs(N_idx, M_idx) = H1_max;
|
||||
fprintf("Test Complete \n\n");
|
||||
end
|
||||
end
|
||||
sgtitle("H_1 (in support set)");
|
||||
|
||||
figure;
|
||||
|
||||
subplot(221);
|
||||
h = heatmap(Ns, Ms, H0_REE_means');
|
||||
h.XLabel = "N";
|
||||
h.YLabel = "M";
|
||||
h.Title = "H_0, mean(REE)";
|
||||
|
||||
subplot(222);
|
||||
h = heatmap(Ns, Ms, H0_REE_maxs');
|
||||
h.XLabel = "N";
|
||||
h.YLabel = "M";
|
||||
h.Title = "H_0, max(REE)";
|
||||
|
||||
subplot(223);
|
||||
h = heatmap(Ns, Ms, H1_REE_means');
|
||||
h.XLabel = "N";
|
||||
h.YLabel = "M";
|
||||
h.Title = "H_1, mean(REE)";
|
||||
|
||||
subplot(224);
|
||||
h = heatmap(Ns, Ms, H1_REE_maxs');
|
||||
h.XLabel = "N";
|
||||
h.YLabel = "M";
|
||||
h.Title = "H_1, max(REE)";
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user