Skip to content

Commit bcfa8ff

Browse files
committed
bench: read variance from solver statistics
1 parent 927c400 commit bcfa8ff

2 files changed

Lines changed: 14 additions & 29 deletions

File tree

benchmarks/knapsack/KnapsackBenchmark.jl

Lines changed: 9 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -251,15 +251,6 @@ function projection_scaling_rows(; cutoff = 1e-10)
251251
end
252252
end
253253

254-
function variance_observer()
255-
variances = Float64[]
256-
callback = function (_mps; variance, kw...)
257-
push!(variances, variance)
258-
return nothing
259-
end
260-
return variances, callback
261-
end
262-
263254
function benchmark_repeated(f; samples)
264255
outputs = Any[]
265256
capture() = begin
@@ -325,10 +316,13 @@ end
325316
function solution_stats(solution)
326317
stats = solution.stats
327318
bonds = stats.max_bonds
319+
variance_index = findlast(variance -> !isnothing(variance), stats.variances)
328320
return (
329321
sweeps = length(stats.energies),
330322
solution_max_bond = isempty(stats.bond_dims) ? 0 : maximum(stats.bond_dims),
331323
solver_elapsed_seconds = isempty(stats.elapsed_times) ? 0.0 : last(stats.elapsed_times),
324+
final_variance =
325+
isnothing(variance_index) ? missing : stats.variances[variance_index],
332326
initial_state_bond = bonds.initial_state,
333327
objective_mpo_bond = bonds.objective,
334328
projection_mpo_bond = isempty(bonds.projections) ? missing : maximum(bonds.projections),
@@ -388,7 +382,7 @@ function best_penalty_sample(instance, model, samples)
388382
return best
389383
end
390384

391-
function solver_options(iterations, cutoff, time_limit, on_iteration)
385+
function solver_options(iterations, cutoff, time_limit)
392386
return (
393387
iterations = iterations,
394388
time_limit = time_limit,
@@ -400,8 +394,6 @@ function solver_options(iterations, cutoff, time_limit, on_iteration)
400394
# Collect variance on every sweep without letting convergence shorten the
401395
# fixed-sweep benchmark workload.
402396
vtol = -Inf,
403-
on_iteration,
404-
callback_every = 1,
405397
verbosity = 0,
406398
)
407399
end
@@ -415,7 +407,6 @@ function benchmark_result_row(
415407
formulation_timed,
416408
case_timed,
417409
solution;
418-
variances,
419410
nvariables,
420411
iterations,
421412
reads,
@@ -476,7 +467,7 @@ function benchmark_result_row(
476467
time_limit_reached = (stats.sweeps < iterations ||
477468
stats.solver_elapsed_seconds > time_limit),
478469
solution_max_bond = stats.solution_max_bond,
479-
final_variance = isempty(variances) ? missing : last(variances),
470+
final_variance = stats.final_variance,
480471
truncation_error = missing,
481472
initial_state_bond = stats.initial_state_bond,
482473
objective_mpo_bond = stats.objective_mpo_bond,
@@ -506,17 +497,15 @@ function projection_row(
506497
end
507498
constraint = formulation_timed.value
508499
case_timed = benchmark_repeated(; samples = timing_samples) do
509-
variances, callback = variance_observer()
510500
Random.seed!(SOLVER_SEED)
511-
options = solver_options(iterations, cutoff, time_limit, callback)
501+
options = solver_options(iterations, cutoff, time_limit)
512502
reported_objective, solution =
513503
TenSolver.maximize(instance.values; constraints = [constraint], options...)
514504
sampling = @timed best_projection_sample(instance, TenSolver.sample(solution, reads))
515505
return (;
516506
reported_objective,
517507
solution,
518508
items = sampling.value,
519-
variances,
520509
sampling = (time = sampling.time, gctime = sampling.gctime, memory = sampling.bytes),
521510
)
522511
end
@@ -531,7 +520,6 @@ function projection_row(
531520
formulation_timed,
532521
case_timed,
533522
output.solution;
534-
variances = output.variances,
535523
nvariables = nitems,
536524
iterations,
537525
reads,
@@ -559,9 +547,8 @@ function penalty_row(
559547
end
560548
model = formulation_timed.value
561549
case_timed = benchmark_repeated(; samples = timing_samples) do
562-
variances, callback = variance_observer()
563550
Random.seed!(SOLVER_SEED)
564-
options = solver_options(iterations, cutoff, time_limit, callback)
551+
options = solver_options(iterations, cutoff, time_limit)
565552
reported_objective, solution =
566553
TenSolver.minimize(model.Q, model.l, model.constant; options...)
567554
sampling = @timed begin
@@ -573,7 +560,6 @@ function penalty_row(
573560
solution,
574561
assignment = sampling.value.assignment,
575562
items = sampling.value.items,
576-
variances,
577563
sampling = (time = sampling.time, gctime = sampling.gctime, memory = sampling.bytes),
578564
)
579565
end
@@ -588,7 +574,6 @@ function penalty_row(
588574
formulation_timed,
589575
case_timed,
590576
output.solution;
591-
variances = output.variances,
592577
nvariables = length(output.assignment),
593578
iterations,
594579
reads,
@@ -628,8 +613,8 @@ not just each solver's encoded objective.
628613
629614
Final-state variance and operator bond dimensions come from solver-reported
630615
statistics.
631-
Truncation error remains unavailable because the callback runs after discarded
632-
singular values have been removed.
616+
Truncation error remains unavailable because solver statistics do not expose
617+
discarded singular values.
633618
634619
When provided, `on_row` is called after each completed row so long runs can
635620
report progress without changing the returned table.

benchmarks/knapsack/README.md

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -115,11 +115,11 @@ imply equal convergence quality.
115115

116116
The projection and effective-Hamiltonian bonds are separate because the latter
117117
drives constrained DMRG cost. The solver reports these bonds through
118-
`solution.stats.max_bonds` and supplies its calculated variance to
119-
`on_iteration`, so the benchmark does not reconstruct the Hamiltonian or repeat
120-
the variance calculation. Truncation error remains empty: the callback runs
121-
after the DMRG sweep has discarded singular values, so that error cannot be
122-
reconstructed from the retained MPS.
118+
`solution.stats.max_bonds` and its checked variances through
119+
`solution.stats.variances`, so the benchmark does not reconstruct the
120+
Hamiltonian or repeat the variance calculation. Truncation error remains empty
121+
because solver statistics do not expose the singular values discarded during a
122+
DMRG sweep.
123123

124124
## Recorded results
125125

0 commit comments

Comments
 (0)