clear all

% number of samples and features
n = 6000;
gamma = 0.5;
p = gamma * n;

% ar-1 covariance
rho = 0.25;
Sigma = toeplitz(rho.^(0:1:p-1));
% Sigma = Sigma;
[W, R] = eig(Sigma);
% reorder the eigenvalues and eigenvectors of Sigma
[~, ind] = sort(diag(R), 'descend');
R = R(ind, ind);
W = W(:, ind);
dR = diag(R);

% ar-1 features
Z = randn(n, p);
% X = Z * sqrtm(Sigma);
% X = Z * W * diag(sqrt(dR)) * W';
% skip Z * W as it has the same distribution as Z
tZ = Z * W;
X = tZ * (sqrt(dR) .* W');
[U, sqrtS, V] = svd(X/sqrt(n));
% S = sqrtS' * sqrtS;
% get the first n entries of S
% dS = diag(S);
% dSn = dS(1:n);
dSn = diag(sqrtS).^2;

sig_energy = 1;
noi_energy = 1;

% signal and response
% beta0 = W(:, 1) / sqrt(dR(1)) * sqrt(sig_energy);
beta0 = W(:, p) / sqrt(dR(p)) * sqrt(sig_energy);
% beta0 = W(:, p) / sqrt(R(p, p));
% beta0 = W(:, n) / sqrt(dR(n)) * sqrt(sig_energy);
y = X * beta0 + sqrt(noi_energy) * randn(n, 1);

% modified tilde quantities
tbeta0 = V' * beta0;
ty = U' * y / sqrt(n);
% tSigma = V' * Sigma *  V;
% split the tSigma into two components
hr_tSigma = sqrt(dR) .* (W' * V);

% setup up lambda grid
% right grid
n_grid_right = 100;
lambda_min = -1 * min(dSn) + 0.01;
lambda_max = 2.5;
lambda_grid_right = linspace(lambda_min, lambda_max, n_grid_right);
% left grid
n_grid_left = 0;
lambda_min = -30;
lambda_max = -10;
lambda_grid_left = linspace(lambda_min, lambda_max, n_grid_left);
% total grid
n_grid = n_grid_left + n_grid_right;
lambda_grid = [lambda_grid_left, lambda_grid_right];

% setup variables to store risk
risk_grid = zeros(n_grid, 1);
gcv_grid = zeros(n_grid, 1);
loocv_grid = zeros(n_grid, 1);

% compute risk
for i = 1:n_grid
    lambda = lambda_grid(i);
    temp_vec_n = sqrt(dSn) ./ (dSn + lambda) .* ty; 
    temp_vec_p = [temp_vec_n; zeros(p-n, 1)];
    temp_vec = temp_vec_p - tbeta0;
    r_temp_vec = hr_tSigma * temp_vec;
    risk = r_temp_vec' * r_temp_vec;
    risk_grid(i) = risk;
    
    gcv_numer = sum(ty .* ty ./ (dSn + lambda).^2);
    gcv_denom = (sum(1 ./ (dSn + lambda)) / n)^2;
    gcv = gcv_numer / gcv_denom;
    gcv_grid(i) = gcv;
    
    loocv_denom_vec = vecnorm((1./sqrt(dSn + lambda)) .* U').^2;
    
    loocv = norm((ty ./ (dSn + lambda)) ./ loocv_denom_vec')^2;
    loocv_grid(i) = loocv;
end

risk_grid = risk_grid + noi_energy;
hold on
grid on
% plot(lambda_grid(1:n_grid_left), risk_grid(1:n_grid_left))
plot(lambda_grid(n_grid_left+1:end), risk_grid(n_grid_left+1:end))
[~, idx] = min(risk_grid);
plot(lambda_grid(idx), risk_grid(idx), '*', 'Color', [0 0.4470 0.7410])
% xline(lambda_grid(idx))
% plot(lambda_grid(1:n_grid_left), gcv_grid(1:n_grid_left))
plot(lambda_grid(n_grid_left+1:end), gcv_grid(n_grid_left+1:end))
[~, idx] = min(gcv_grid);
% xline(lambda_grid(idx))
plot(lambda_grid(idx), gcv_grid(idx), 'o', 'Color', [0.8500 0.3250 0.0980])
plot(lambda_grid(n_grid_left+1:end), loocv_grid(n_grid_left+1:end))
[~, idx] = min(loocv_grid);
% xline(lambda_grid(idx))
plot(lambda_grid(idx), loocv_grid(idx), '+', 'Color', [0.9290 0.6940 0.1250])
legend('risk', 'min risk', 'gcv', 'min gcv', 'loocv', 'min loocv', 'Location', 'best')
% legend('risk left', 'risk right', 'min risk', 'gcv left', 'gcv right', 'min gcv', 'Location', 'best')
xlabel('\lambda')
