-
Notifications
You must be signed in to change notification settings - Fork 5
Expand file tree
/
Copy pathfistaN.m
More file actions
116 lines (99 loc) · 2.62 KB
/
Copy pathfistaN.m
File metadata and controls
116 lines (99 loc) · 2.62 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
function [x lambda errvec] = fistaN(AA,Ab,sparsity,tol,maxit,Q,svals)
% [x lambda errvec] = fista(AA,Ab,sparsity,tol,maxit,Q)
%
% Solves the following problem via FISTA (normal eq):
%
% minimize (1/2)||AAx-Ab||_2^2 + lambda*||Qx||_1
%
% -AA is an [n x n] matrix (or function handle)
% -sparsity is the fraction of zeros (0.1=10% zeros)
% -tol/maxit are tolerance and max. no. iterations
% -Q is a wavelet transform (Q*x and Q'*x - see HWT)
% -svals is [largest smallest] singular values of AA
%
% -lambda that yields the required sparsity (scalar)
% -errvec is the rel. change at each iteration (vector)
%
%% check arguments
if nargin<3 || nargin>7
error('Wrong number of input arguments');
end
if nargin<4 || isempty(tol)
tol = 1e-3;
end
if nargin<5 || isempty(maxit)
maxit = 100;
end
if nargin<6 || isempty(Q)
Q = 1;
end
if nargin<7 || isempty(svals)
svals = [];
else
normAA = svals(1);
end
if isnumeric(AA)
AA = @(arg) AA*arg;
end
if ~iscolumn(Ab)
error('Ab must be a column vector');
end
if ~isscalar(sparsity) || sparsity<0 || sparsity>1
error('sparsity must be a scalar between 0 and 1.');
end
%% solve by FISTA
time = tic();
z = Ab;
t = 1;
for iter = 1:maxit
Az = AA(z);
% steepest descent along Ab
if iter==1
alpha = (z'*z) / (Az'*Az);
z = alpha * z;
x = z;
end
% power iteration to get norm(AA)
if ~exist('normAA','var')
tmp = Az/norm(Az);
for k = 1:20
if k>1
tmp = tmp/normAA(k-1);
end
tmp = AA(tmp);
normAA(k) = norm(tmp);
end
normAA = norm(tmp);
end
z = z + (Ab - Az) / normAA;
xold = x;
x = Q*z;
[x lambda(iter)] = shrinkage(x,sparsity);
x = Q'*x;
errvec(iter) = norm(x-xold)/norm(x);
if errvec(iter) < tol
break;
end
% FISTA-ADA
%if numel(svals)==2
% r = 4 * (1-sqrt(prod(svals)))^2 / abs(1-prod(svals));
%else
r = 4;
%end
t0 = t;
t = (1+sqrt(1+r*t^2))/2;
z = x + ((t0-1)/t) * (x-xold);
end
% report convergence
if iter < maxit
fprintf('%s converged at iteration %i to a solution with relative error %.1e. ',mfilename,iter,errvec(iter));
else
fprintf('%s stopped at iteration %i with relative error %.1e without converging to the desired tolerance %.1e. ',mfilename,iter,errvec(iter),tol);
end
toc(time);
%% shrink based on sparsity => return lambda
function [z lambda] = shrinkage(z,sparsity)
absz = abs(z);
v = sort(absz,'ascend');
lambda = interp1(v,sparsity*numel(z),'linear',0);
z = sign(z) .* max(absz-lambda, 0); % complex ok