@@ -38,7 +38,7 @@ def __init__(
3838 mult = 1.5 ,
3939 S2_grid = 441 ,
4040 Nstep_yI = 10 ,
41- perform_pr = False ,
41+ pr_iters = None ,
4242 verbose = True ,
4343 ** kwargs ,
4444 ):
@@ -62,7 +62,8 @@ def __init__(
6262 :param mult: Step-size multiplier for the ADMM primal update.
6363 :param S2_grid: Number of sphere samples used to discretize SO(3).
6464 :param Nstep_yI: Number of inequality-multiplier updates per ADMM iteration.
65- :param perform_pr: Whether to apply proximal refinement after ADMM.
65+ :param pr_iters: Number of proximal refinement iterations. Default of None
66+ does not perform proximal refinement. Recommended value is 4.
6667 :param verbose: Whether to log ADMM progress.
6768 """
6869
@@ -85,7 +86,6 @@ def __init__(
8586 self .mult = mult
8687 self .S2_grid = S2_grid
8788 self .Nstep_yI = Nstep_yI
88- self .perform_pr = perform_pr
8989 self .verbose = verbose
9090
9191 # Handle symmetry
@@ -110,6 +110,13 @@ def __init__(
110110 self .sym_euler = self .sym_grp .rotations .angles
111111 self .n_sym = len (self .sym_euler )
112112
113+ # Set up proximal refinement terms
114+ if pr_iters is not None :
115+ self .pr_weights = 1 / (1 + np .arange (self .Lmax ))
116+ self .pr_penalty = [1 ] * pr_iters
117+ self .pr_rank = list (range (pr_iters - 1 , - 1 , - 1 )) # [pr_iters - 1,..., 0]
118+ self .pr_iters = pr_iters
119+
113120 self ._build_full_pft ()
114121
115122 def _build_full_pft (self ):
@@ -276,15 +283,12 @@ def perform_admm(self):
276283 """
277284 X_est = self .admm_sym_J (self .C , self .verbose )
278285
279- if self .perform_pr :
280- weight = 1 / (1 + np .arange (self .Lmax ))
281- Penalty = [1 , 1 , 1 , 1 ]
282- r = [3 , 2 , 1 , 0 ]
286+ if self .pr_iters is not None :
283287 X_est = self .proximal_refine (
284288 X_est ,
285- weight ,
286- Penalty ,
287- r ,
289+ self . pr_weights ,
290+ self . pr_penalty ,
291+ self . pr_rank ,
288292 )
289293 self .X_est = X_est
290294
@@ -980,11 +984,10 @@ def low_rank_proj(X, r_step):
980984 Xproj .append (tmp )
981985 return Xproj
982986
983- Niter = len (r )
984987 CC = [None ] * self .Lmax
985988 current = [np .copy (Xk ) for Xk in X_admm ]
986989
987- for step in range (Niter ):
990+ for step in range (self . pr_iters ):
988991 X_proj = low_rank_proj (current , r [step ])
989992
990993 for k in range (self .Lmax ):
@@ -999,7 +1002,7 @@ def low_rank_proj(X, r_step):
9991002 logger .info (
10001003 "Proximal refine step %d/%d: relative update %.3e" ,
10011004 step + 1 ,
1002- Niter ,
1005+ self . pr_iters ,
10031006 rel_change (X_next , current ),
10041007 )
10051008
0 commit comments