Skip to content

Commit 0832fe9

Browse files
Add shortcut
1 parent 3421dd3 commit 0832fe9

4 files changed

Lines changed: 55 additions & 37 deletions

File tree

src/bell_frank_wolfe.jl

Lines changed: 6 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@ function bell_frank_wolfe(
3939
v0 = one(T),
4040
epsilon = 10Base.rtoldefault(T),
4141
verbose = 0,
42+
verbose_init = verbose > 0,
4243
shr2 = NaN,
4344
mode::Int = 0,
4445
nb::Int = 10^2,
@@ -54,6 +55,7 @@ function bell_frank_wolfe(
5455
reset_dots_A::Bool = false, # warm start (technical)
5556
reset_dots_b::Bool = true, # warm start (technical)
5657
lazy::Bool = true, # default in FW package is false
58+
shortcut = 0, # early termination criterion (outside only)
5759
max_iteration::Int = 10^9, # default in FW package is 10^4
5860
nb_increment_interval::Int = 10^4,
5961
callback_interval::Int = verbose > 0 ? 10^4 : typemax(Int),
@@ -66,7 +68,7 @@ function bell_frank_wolfe(
6668
kwargs...,
6769
) where {T <: Number, N}
6870
Random.seed!(seed)
69-
LMO, DS, m, o, sym, deflate, inflate = _bfw_init(p, v0, prob, marg, o, sym, deflate, inflate, verbose)
71+
LMO, DS, m, o, sym, deflate, inflate = _bfw_init(p, v0, prob, marg, o, sym, deflate, inflate, verbose_init)
7072
if verbose > 0
7173
println("Visibility: ", v0)
7274
end
@@ -111,16 +113,14 @@ function bell_frank_wolfe(
111113
println("Active set initialised")
112114
end
113115
end
114-
if verbose > 0
115-
println()
116-
end
117116
callback = build_callback(
118117
rp,
119118
v0,
120119
ro,
121120
shr2 ^ (prob ? (N ÷ 2) / 2 : N / 2),
122121
verbose,
123122
epsilon,
123+
shortcut,
124124
nb_increment_interval,
125125
callback_interval,
126126
hyperplane_interval,
@@ -149,8 +149,7 @@ function bell_frank_wolfe(
149149
primal = res.primal
150150
dual_gap = res.dual_gap
151151
as = res.active_set
152-
if verbose 2
153-
println()
152+
if verbose == 2
154153
@printf("Primal: %.2e\n", primal)
155154
@printf("FW gap: %.2e\n", dual_gap)
156155
@printf("#Atoms: %d\n", length(as))
@@ -196,14 +195,14 @@ function bell_frank_wolfe(
196195
if verbose > 0
197196
if verbose 2 && mode_last 0
198197
@printf("FW gap: %.2e\n", dual_gap) # recomputed FW gap (usually with a more reliable heuristic)
199-
println()
200198
end
201199
if primal > dual_gap
202200
@printf("v_c ≤ %f\n", β)
203201
elseif !isnan(shr2)
204202
ν = 1 / (1 + norm(vp - as.x, 2))
205203
@printf("v_c ≥ %f (%f)\n", shr2^(N / 2) * ν * v0, shr2^(N / 2) * v0)
206204
end
205+
println()
207206
end
208207
if save
209208
serialize(file * ".dat", ActiveSetStorage(as))

src/callback.jl

Lines changed: 44 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ function build_callback(
55
shr2,
66
verbose,
77
epsilon,
8+
shortcut,
89
nb_increment_interval,
910
callback_interval,
1011
hyperplane_interval,
@@ -37,34 +38,14 @@ function build_callback(
3738
verbose_hyperplane = (verbose || save) && hyperplane_interval != typemax(Int)
3839
verbose_bound = verbose && bound_interval != typemax(Int)
3940
if verbose
40-
@printf(
41-
stdout,
42-
"%s %s %s %s %s %s %s\n",
43-
lpad("Iteration", 12),
44-
lpad("Primal", 10),
45-
lpad("Dual gap", 10),
46-
lpad("Time (sec)", 10),
47-
lpad("#It/sec", 10),
48-
lpad("#Atoms", 7),
49-
lpad("#LMO", 7)
50-
)
41+
_print_headers()
5142
end
5243
function callback(state, active_set, args...)
5344
if mod(state.t, nb_increment_interval) == 0
5445
state.lmo.lmo.nb += 1
5546
end
5647
if verbose && mod(state.t, callback_interval) == 0
57-
@printf(
58-
stdout,
59-
"%s %.4e %.4e %.4e %.4e %s %s\n",
60-
lpad(state.t, 12),
61-
state.primal,
62-
state.dual_gap,
63-
state.time,
64-
state.t / state.time,
65-
lpad(length(active_set), 7),
66-
lpad(state.lmo.lmo.cnt, 7)
67-
)
48+
_print_callback(state.t, state, active_set)
6849
end
6950
if verbose_hyperplane && mod(state.t, hyperplane_interval) == 0
7051
a = -state.gradient # v*p+(1-v)*o-active_set.x
@@ -83,10 +64,47 @@ function build_callback(
8364
if save && mod(state.t, save_interval) == 0
8465
serialize(file * "_tmp.dat", ActiveSetStorage(active_set))
8566
end
86-
# if state.dual_gap < state.primal / 2
87-
# return false
88-
# end
89-
return state.primal > epsilon
67+
if shortcut > 0 && state.dual_gap < state.primal / shortcut
68+
if verbose
69+
_print_callback("Shortcut", state, active_set)
70+
end
71+
return false
72+
end
73+
if state.primal epsilon || state.dual_gap epsilon
74+
if verbose
75+
_print_callback("Last", state, active_set)
76+
end
77+
return false
78+
end
79+
return true
9080
end
9181
return callback
9282
end
83+
84+
function _print_headers()
85+
@printf(
86+
stdout,
87+
"%s %s %s %s %s %s %s\n",
88+
lpad("Iteration", 12),
89+
lpad("Primal", 10),
90+
lpad("Dual gap", 10),
91+
lpad("Time (sec)", 10),
92+
lpad("#It/sec", 10),
93+
lpad("#Atoms", 7),
94+
lpad("#LMO", 7)
95+
)
96+
end
97+
98+
function _print_callback(it, state, active_set)
99+
@printf(
100+
stdout,
101+
"%s %.4e %.4e %.4e %.4e %s %s\n",
102+
lpad(it, 12),
103+
state.primal,
104+
state.dual_gap,
105+
state.time,
106+
state.t / state.time,
107+
lpad(length(active_set), 7),
108+
lpad(state.lmo.lmo.cnt, 7)
109+
)
110+
end

src/nonlocality_threshold.jl

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,17 +25,18 @@ function nonlocality_threshold(
2525
deflate = identity,
2626
inflate = identity,
2727
verbose = 0,
28+
shortcut = 100,
2829
kwargs...,
2930
) where {T <: Number, N}
30-
_, _, _, o, sym, deflate, inflate = _bfw_init(p, 0, prob, marg, o, sym, deflate, inflate, verbose)
31+
_, _, _, o, sym, deflate, inflate = _bfw_init(p, 0, prob, marg, o, sym, deflate, inflate, verbose > 0)
3132
lower_bound = zero(T)
3233
upper_bound = one(T)
3334
local_model = nothing
3435
bell_inequality = nothing
3536
while upper_bound - lower_bound > 10.0^(-precision)
36-
res = bell_frank_wolfe(p; v0, epsilon, marg, o, sym, deflate, inflate, kwargs...)
37+
res = bell_frank_wolfe(p; v0, epsilon, marg, o, sym, deflate, inflate, verbose, verbose_init = false, shortcut, kwargs...)
3738
x, ds, primal, dual_gap, as, M, β = res
38-
if primal > 10epsilon && dual_gap > 10epsilon
39+
if dual_gap * shortcut primal && primal > 10epsilon && dual_gap > 10epsilon
3940
@warn "Please increase nb or max_iteration"
4041
end
4142
if dual_gap < primal

src/utils.jl

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1151,7 +1151,7 @@ function _bfw_init(p::Array{T, N}, v0, prob, marg, o, sym, deflate, inflate, ver
11511151
sym = false
11521152
end
11531153
end
1154-
if verbose > 1
1154+
if verbose
11551155
println(" #Inputs: ", all(diff(m) .== 0) ? m[end] - (marg && !prob) : m .- (marg && !prob))
11561156
println(" Symmetric: ", sym)
11571157
println(" Dimension: ", length(deflate(p)))

0 commit comments

Comments
 (0)