1use crate::types::Number;
7use crate::utils::{cpu_time, sys_time, wallclock_time};
8use std::cell::{Cell, RefCell};
9use std::rc::Rc;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum DeadlineKind {
14 Wall,
16 Cpu,
18}
19
20#[derive(Debug, Clone)]
42pub struct Deadline {
43 inner: Rc<DeadlineInner>,
44}
45
46#[derive(Debug)]
47struct DeadlineInner {
48 wall_start: Number,
49 cpu_start: Number,
50 max_wall: Number,
51 max_cpu: Number,
52}
53
54impl Deadline {
55 pub fn new(max_wall: Number, max_cpu: Number) -> Self {
60 Self {
61 inner: Rc::new(DeadlineInner {
62 wall_start: wallclock_time(),
63 cpu_start: cpu_time(),
64 max_wall,
65 max_cpu,
66 }),
67 }
68 }
69
70 pub fn exceeded(&self) -> Option<DeadlineKind> {
77 if cpu_time() - self.inner.cpu_start >= self.inner.max_cpu {
78 return Some(DeadlineKind::Cpu);
79 }
80 if wallclock_time() - self.inner.wall_start >= self.inner.max_wall {
81 return Some(DeadlineKind::Wall);
82 }
83 None
84 }
85
86 pub fn max_wall(&self) -> Number {
88 self.inner.max_wall
89 }
90
91 pub fn max_cpu(&self) -> Number {
93 self.inner.max_cpu
94 }
95
96 pub fn remaining_wall(&self) -> Number {
103 self.inner.max_wall - (wallclock_time() - self.inner.wall_start)
104 }
105
106 pub fn remaining_cpu(&self) -> Number {
108 self.inner.max_cpu - (cpu_time() - self.inner.cpu_start)
109 }
110}
111
112#[derive(Debug)]
116pub struct TimedTask {
117 enabled: Cell<bool>,
118 start_called: Cell<bool>,
119 end_called: Cell<bool>,
120 start_cpu: Cell<Number>,
121 start_sys: Cell<Number>,
122 start_wall: Cell<Number>,
123 total_cpu: Cell<Number>,
124 total_sys: Cell<Number>,
125 total_wall: Cell<Number>,
126}
127
128impl Default for TimedTask {
129 fn default() -> Self {
130 Self {
131 enabled: Cell::new(true),
132 start_called: Cell::new(false),
133 end_called: Cell::new(true),
134 start_cpu: Cell::new(0.0),
135 start_sys: Cell::new(0.0),
136 start_wall: Cell::new(0.0),
137 total_cpu: Cell::new(0.0),
138 total_sys: Cell::new(0.0),
139 total_wall: Cell::new(0.0),
140 }
141 }
142}
143
144impl TimedTask {
145 pub fn new() -> Self {
146 Self::default()
147 }
148
149 pub fn enable(&self) {
150 self.enabled.set(true);
151 }
152 pub fn disable(&self) {
153 self.enabled.set(false);
154 }
155 pub fn is_enabled(&self) -> bool {
156 self.enabled.get()
157 }
158 pub fn is_started(&self) -> bool {
159 self.start_called.get()
160 }
161
162 pub fn reset(&self) {
163 self.total_cpu.set(0.0);
164 self.total_sys.set(0.0);
165 self.total_wall.set(0.0);
166 self.start_called.set(false);
167 self.end_called.set(true);
168 }
169
170 pub fn start(&self) {
171 if !self.enabled.get() {
172 return;
173 }
174 self.end_called.set(false);
175 self.start_called.set(true);
176 self.start_cpu.set(cpu_time());
177 self.start_sys.set(sys_time());
178 self.start_wall.set(wallclock_time());
179 }
180
181 pub fn end(&self) {
182 if !self.enabled.get() {
183 return;
184 }
185 self.end_called.set(true);
186 self.start_called.set(false);
187 self.total_cpu
188 .set(self.total_cpu.get() + cpu_time() - self.start_cpu.get());
189 self.total_sys
190 .set(self.total_sys.get() + sys_time() - self.start_sys.get());
191 self.total_wall
192 .set(self.total_wall.get() + wallclock_time() - self.start_wall.get());
193 }
194
195 pub fn end_if_started(&self) {
196 if !self.enabled.get() {
197 return;
198 }
199 if self.start_called.get() {
200 self.end();
201 }
202 }
203
204 pub fn total_cpu_time(&self) -> Number {
205 self.total_cpu.get()
206 }
207 pub fn total_sys_time(&self) -> Number {
208 self.total_sys.get()
209 }
210 pub fn total_wallclock_time(&self) -> Number {
211 self.total_wall.get()
212 }
213
214 pub fn live_wallclock_time(&self) -> Number {
220 if self.enabled.get() && self.start_called.get() {
221 self.total_wall.get() + wallclock_time() - self.start_wall.get()
222 } else {
223 self.total_wall.get()
224 }
225 }
226
227 pub fn live_cpu_time(&self) -> Number {
230 if self.enabled.get() && self.start_called.get() {
231 self.total_cpu.get() + cpu_time() - self.start_cpu.get()
232 } else {
233 self.total_cpu.get()
234 }
235 }
236
237 pub fn guard(&self) -> TimedGuard<'_> {
241 self.start();
242 TimedGuard { task: Some(self) }
243 }
244}
245
246#[must_use = "the guard ends the timer when dropped; bind it to a variable"]
250pub struct TimedGuard<'a> {
251 task: Option<&'a TimedTask>,
252}
253
254impl<'a> TimedGuard<'a> {
255 pub fn stop(mut self) {
259 if let Some(t) = self.task.take() {
260 t.end_if_started();
261 }
262 }
263}
264
265impl<'a> Drop for TimedGuard<'a> {
266 fn drop(&mut self) {
267 if let Some(t) = self.task.take() {
268 t.end_if_started();
269 }
270 }
271}
272
273#[derive(Debug, Default)]
279pub struct TimingStatistics {
280 pub overall_alg: TimedTask,
281 pub print_problem_statistics: TimedTask,
282 pub initialize_iterates: TimedTask,
283 pub update_hessian: TimedTask,
284 pub output_iteration: TimedTask,
285 pub update_barrier_parameter: TimedTask,
286 pub compute_search_direction: TimedTask,
287 pub compute_acceptable_trial_point: TimedTask,
288 pub accept_trial_point: TimedTask,
289 pub check_convergence: TimedTask,
290 pub fire_intermediate: TimedTask,
298
299 pub linear_system_symbolic_factorization: TimedTask,
300 pub linear_system_factorization: TimedTask,
301 pub linear_system_back_solve: TimedTask,
302 pub quality_function_search: TimedTask,
303 pub total_callback_time: TimedTask,
304 pub total_function_evaluation_time: TimedTask,
305 pub eval_obj: TimedTask,
306 pub eval_grad_obj: TimedTask,
307 pub eval_constr: TimedTask,
308 pub eval_constr_jac: TimedTask,
309 pub eval_lag_hess: TimedTask,
310}
311
312impl TimingStatistics {
313 pub fn new() -> Self {
314 Self::default()
315 }
316
317 pub fn report(&self) -> String {
325 use std::fmt::Write as _;
326 let mut s = String::new();
327 let row = |s: &mut String, label: &str, t: &TimedTask| {
328 let _ = writeln!(
329 s,
330 "{label:<42} {wall:>10.3}s",
331 wall = t.total_wallclock_time()
332 );
333 };
334 s.push_str("\nTiming Statistics:\n");
335 row(
336 &mut s,
337 "OverallAlgorithm....................:",
338 &self.overall_alg,
339 );
340 row(
341 &mut s,
342 " InitializeIterates.................:",
343 &self.initialize_iterates,
344 );
345 row(
346 &mut s,
347 " UpdateHessian......................:",
348 &self.update_hessian,
349 );
350 row(
351 &mut s,
352 " OutputIteration....................:",
353 &self.output_iteration,
354 );
355 row(
356 &mut s,
357 " UpdateBarrierParameter.............:",
358 &self.update_barrier_parameter,
359 );
360 row(
361 &mut s,
362 " ComputeSearchDirection.............:",
363 &self.compute_search_direction,
364 );
365 row(
366 &mut s,
367 " ComputeAcceptableTrialPoint........:",
368 &self.compute_acceptable_trial_point,
369 );
370 row(
371 &mut s,
372 " AcceptTrialPoint...................:",
373 &self.accept_trial_point,
374 );
375 row(
376 &mut s,
377 " CheckConvergence...................:",
378 &self.check_convergence,
379 );
380 row(
381 &mut s,
382 " FireIntermediateCallback...........:",
383 &self.fire_intermediate,
384 );
385 row(
386 &mut s,
387 "LinearSystemSymbolicFactorization...:",
388 &self.linear_system_symbolic_factorization,
389 );
390 row(
391 &mut s,
392 "LinearSystemFactorization...........:",
393 &self.linear_system_factorization,
394 );
395 row(
396 &mut s,
397 "LinearSystemBackSolve...............:",
398 &self.linear_system_back_solve,
399 );
400 row(
401 &mut s,
402 "QualityFunctionSearch...............:",
403 &self.quality_function_search,
404 );
405 row(
406 &mut s,
407 "TotalFunctionEvaluations............:",
408 &self.total_function_evaluation_time,
409 );
410 row(
411 &mut s,
412 " ObjectiveFunctionEvaluations.......:",
413 &self.eval_obj,
414 );
415 row(
416 &mut s,
417 " ObjectiveGradientEvaluations.......:",
418 &self.eval_grad_obj,
419 );
420 row(
421 &mut s,
422 " ConstraintEvaluations..............:",
423 &self.eval_constr,
424 );
425 row(
426 &mut s,
427 " ConstraintJacobianEvaluations......:",
428 &self.eval_constr_jac,
429 );
430 row(
431 &mut s,
432 " LagrangianHessianEvaluations.......:",
433 &self.eval_lag_hess,
434 );
435 s
436 }
437
438 pub fn wall_time_breakdown(&self) -> Vec<(&'static str, Number)> {
455 let symbolic = self
456 .linear_system_symbolic_factorization
457 .total_wallclock_time();
458 let factorization = self.linear_system_factorization.total_wallclock_time();
459 let back_solve = self.linear_system_back_solve.total_wallclock_time();
460 vec![
461 ("overall_alg", self.overall_alg.total_wallclock_time()),
462 ("update_hessian", self.update_hessian.total_wallclock_time()),
463 (
464 "compute_search_direction",
465 self.compute_search_direction.total_wallclock_time(),
466 ),
467 ("linear_system_total", symbolic + factorization + back_solve),
468 ("linear_system_symbolic_factorization", symbolic),
469 ("linear_system_factorization", factorization),
470 ("linear_system_back_solve", back_solve),
471 (
472 "function_evaluations_total",
473 self.total_function_evaluation_time.total_wallclock_time(),
474 ),
475 ("eval_objective", self.eval_obj.total_wallclock_time()),
476 ("eval_gradient", self.eval_grad_obj.total_wallclock_time()),
477 ("eval_constraints", self.eval_constr.total_wallclock_time()),
478 (
479 "eval_constraint_jacobian",
480 self.eval_constr_jac.total_wallclock_time(),
481 ),
482 (
483 "eval_lagrangian_hessian",
484 self.eval_lag_hess.total_wallclock_time(),
485 ),
486 (
487 "total_callback",
488 self.total_callback_time.total_wallclock_time(),
489 ),
490 (
491 "fire_intermediate",
492 self.fire_intermediate.total_wallclock_time(),
493 ),
494 ]
495 }
496
497 pub fn set_detailed_enabled(&self, on: bool) {
514 let set = |t: &TimedTask| {
515 if on {
516 t.enable();
517 } else {
518 t.disable();
519 }
520 };
521 set(&self.print_problem_statistics);
523 set(&self.initialize_iterates);
524 set(&self.update_hessian);
525 set(&self.output_iteration);
526 set(&self.update_barrier_parameter);
527 set(&self.compute_search_direction);
528 set(&self.compute_acceptable_trial_point);
529 set(&self.accept_trial_point);
530 set(&self.check_convergence);
531 set(&self.fire_intermediate);
532 set(&self.linear_system_symbolic_factorization);
533 set(&self.linear_system_factorization);
534 set(&self.linear_system_back_solve);
535 set(&self.quality_function_search);
536 set(&self.total_callback_time);
537 set(&self.total_function_evaluation_time);
538 set(&self.eval_obj);
539 set(&self.eval_grad_obj);
540 set(&self.eval_constr);
541 set(&self.eval_constr_jac);
542 set(&self.eval_lag_hess);
543 }
544
545 pub fn reset(&self) {
547 self.overall_alg.reset();
548 self.print_problem_statistics.reset();
549 self.initialize_iterates.reset();
550 self.update_hessian.reset();
551 self.output_iteration.reset();
552 self.update_barrier_parameter.reset();
553 self.compute_search_direction.reset();
554 self.compute_acceptable_trial_point.reset();
555 self.accept_trial_point.reset();
556 self.check_convergence.reset();
557 self.fire_intermediate.reset();
558 self.linear_system_symbolic_factorization.reset();
559 self.linear_system_factorization.reset();
560 self.linear_system_back_solve.reset();
561 self.quality_function_search.reset();
562 self.total_callback_time.reset();
563 self.total_function_evaluation_time.reset();
564 self.eval_obj.reset();
565 self.eval_grad_obj.reset();
566 self.eval_constr.reset();
567 self.eval_constr_jac.reset();
568 self.eval_lag_hess.reset();
569 }
570}
571
572#[derive(Debug, Clone, Copy, PartialEq, Eq)]
577pub enum LinearSystemPhase {
578 SymbolicFactorization,
580 Factorization,
582 BackSolve,
584}
585
586#[derive(Debug, Default)]
603pub struct ConvexTimingStatistics {
604 pub overall_alg: TimedTask,
608 pub extraction: TimedTask,
610 pub presolve: TimedTask,
612 pub solve: TimedTask,
614 pub solution_recovery: TimedTask,
617
618 pub linear_system_symbolic_factorization: TimedTask,
619 pub linear_system_factorization: TimedTask,
620 pub linear_system_back_solve: TimedTask,
621}
622
623impl ConvexTimingStatistics {
624 pub fn new() -> Self {
625 Self::default()
626 }
627
628 pub fn report(&self) -> String {
633 use std::fmt::Write as _;
634 let mut s = String::new();
635 let row = |s: &mut String, label: &str, t: &TimedTask| {
636 let _ = writeln!(
637 s,
638 "{label:<42} {wall:>10.3}s",
639 wall = t.total_wallclock_time()
640 );
641 };
642 s.push_str("\nTiming Statistics:\n");
643 row(
644 &mut s,
645 "OverallAlgorithm....................:",
646 &self.overall_alg,
647 );
648 row(
649 &mut s,
650 " ProblemExtraction..................:",
651 &self.extraction,
652 );
653 row(
654 &mut s,
655 " Presolve...........................:",
656 &self.presolve,
657 );
658 row(&mut s, " ConvexSolve........................:", &self.solve);
659 row(
660 &mut s,
661 " SolutionRecovery...................:",
662 &self.solution_recovery,
663 );
664 row(
665 &mut s,
666 "LinearSystemSymbolicFactorization...:",
667 &self.linear_system_symbolic_factorization,
668 );
669 row(
670 &mut s,
671 "LinearSystemFactorization...........:",
672 &self.linear_system_factorization,
673 );
674 row(
675 &mut s,
676 "LinearSystemBackSolve...............:",
677 &self.linear_system_back_solve,
678 );
679 s
680 }
681
682 pub fn set_detailed_enabled(&self, on: bool) {
692 let set = |t: &TimedTask| {
693 if on {
694 t.enable();
695 } else {
696 t.disable();
697 }
698 };
699 set(&self.extraction);
700 set(&self.presolve);
701 set(&self.solve);
702 set(&self.solution_recovery);
703 set(&self.linear_system_symbolic_factorization);
704 set(&self.linear_system_factorization);
705 set(&self.linear_system_back_solve);
706 }
707
708 fn phase(&self, phase: LinearSystemPhase) -> &TimedTask {
709 match phase {
710 LinearSystemPhase::SymbolicFactorization => &self.linear_system_symbolic_factorization,
711 LinearSystemPhase::Factorization => &self.linear_system_factorization,
712 LinearSystemPhase::BackSolve => &self.linear_system_back_solve,
713 }
714 }
715}
716
717thread_local! {
718 static CONVEX_TIMING: RefCell<Option<Rc<ConvexTimingStatistics>>> =
724 const { RefCell::new(None) };
725}
726
727#[must_use = "the scope ends when the guard is dropped; bind it to a variable"]
744pub struct ConvexTimingScope {
745 previous: Option<Rc<ConvexTimingStatistics>>,
746}
747
748impl ConvexTimingScope {
749 pub fn open(stats: &Rc<ConvexTimingStatistics>) -> Self {
750 let previous = CONVEX_TIMING.with(|slot| slot.borrow_mut().replace(Rc::clone(stats)));
751 Self { previous }
752 }
753}
754
755impl Drop for ConvexTimingScope {
756 fn drop(&mut self) {
757 CONVEX_TIMING.with(|slot| *slot.borrow_mut() = self.previous.take());
758 }
759}
760
761pub fn time_linear_system<T>(phase: LinearSystemPhase, f: impl FnOnce() -> T) -> T {
767 let Some(stats) = CONVEX_TIMING.with(|slot| slot.borrow().clone()) else {
768 return f();
769 };
770 let task = stats.phase(phase);
771 let _guard = task.guard();
772 f()
773}
774
775#[cfg(test)]
776mod tests {
777 use super::*;
778
779 #[test]
780 fn deadline_unbounded_never_trips() {
781 let d = Deadline::new(1e6, 1e6);
784 assert!(d.exceeded().is_none());
785 assert_eq!(d.max_wall(), 1e6);
786 assert_eq!(d.max_cpu(), 1e6);
787 }
788
789 #[test]
790 fn deadline_zero_wall_trips_wall() {
791 let d = Deadline::new(0.0, 1e6);
795 for _ in 0..10_000 {
799 if d.exceeded().is_some() {
800 break;
801 }
802 std::hint::black_box(0u64);
803 }
804 assert_eq!(d.exceeded(), Some(DeadlineKind::Wall));
805 }
806
807 #[test]
808 fn deadline_zero_cpu_takes_priority() {
809 let d = Deadline::new(0.0, 0.0);
813 for _ in 0..10_000 {
814 if d.exceeded().is_some() {
815 break;
816 }
817 std::hint::black_box(0u64);
818 }
819 assert_eq!(d.exceeded(), Some(DeadlineKind::Cpu));
820 }
821
822 #[test]
823 fn deadline_remaining_reports_budget_left() {
824 let d = Deadline::new(1e6, 1e6);
828 let rw = d.remaining_wall();
829 let rc = d.remaining_cpu();
830 assert!(rw > 0.0 && rw <= 1e6, "remaining_wall out of range: {rw}");
831 assert!(rc > 0.0 && rc <= 1e6, "remaining_cpu out of range: {rc}");
832 }
833
834 #[test]
835 fn deadline_remaining_goes_nonpositive_once_crossed() {
836 let d = Deadline::new(0.0, 1e6);
839 for _ in 0..10_000 {
840 if d.exceeded().is_some() {
841 break;
842 }
843 std::hint::black_box(0u64);
844 }
845 assert!(
846 d.remaining_wall() <= 0.0,
847 "remaining_wall should be non-positive once the budget is crossed"
848 );
849 }
850
851 #[test]
852 fn deadline_is_cheaply_clonable_and_shares_start() {
853 let d = Deadline::new(1e6, 1e6);
857 let d2 = d.clone();
858 assert_eq!(d2.max_wall(), d.max_wall());
859 assert_eq!(d2.max_cpu(), d.max_cpu());
860 assert!(d2.exceeded().is_none());
861 }
862
863 #[test]
864 fn start_end_accumulates_nonneg() {
865 let t = TimedTask::new();
866 t.start();
867 for _ in 0..1000 {
868 std::hint::black_box(0u64);
869 }
870 t.end();
871 assert!(t.total_wallclock_time() >= 0.0);
872 }
873
874 #[test]
875 fn disabled_is_noop() {
876 let t = TimedTask::new();
877 t.disable();
878 t.start();
879 t.end();
880 assert_eq!(t.total_wallclock_time(), 0.0);
881 }
882
883 #[test]
884 fn set_detailed_enabled_gates_all_but_overall_alg() {
885 let stats = TimingStatistics::new();
886 assert!(stats.overall_alg.is_enabled());
888 assert!(stats.eval_obj.is_enabled());
889 assert!(stats.check_convergence.is_enabled());
890
891 stats.set_detailed_enabled(false);
896 assert!(stats.overall_alg.is_enabled(), "overall_alg must stay live");
897 assert!(!stats.eval_obj.is_enabled());
898 assert!(!stats.check_convergence.is_enabled());
899 assert!(!stats.total_function_evaluation_time.is_enabled());
900 assert!(!stats.linear_system_factorization.is_enabled());
901
902 stats.eval_obj.start();
904 stats.eval_obj.end();
905 assert_eq!(stats.eval_obj.total_wallclock_time(), 0.0);
906
907 stats.set_detailed_enabled(true);
909 assert!(stats.eval_obj.is_enabled());
910 assert!(stats.check_convergence.is_enabled());
911 }
912
913 #[test]
914 fn end_if_started_handles_unstarted() {
915 let t = TimedTask::new();
916 t.end_if_started();
917 assert_eq!(t.total_wallclock_time(), 0.0);
918 }
919
920 #[test]
921 fn wall_time_breakdown_reports_subsystems() {
922 let stats = TimingStatistics::new();
923 stats.linear_system_factorization.start();
927 stats.linear_system_factorization.end();
928 stats.eval_lag_hess.start();
929 stats.eval_lag_hess.end();
930
931 let bd = stats.wall_time_breakdown();
932 let get = |k: &str| bd.iter().find(|(label, _)| *label == k).map(|(_, v)| *v);
933
934 for key in [
936 "overall_alg",
937 "linear_system_total",
938 "linear_system_factorization",
939 "linear_system_back_solve",
940 "function_evaluations_total",
941 "eval_objective",
942 "eval_gradient",
943 "eval_constraints",
944 "eval_constraint_jacobian",
945 "eval_lagrangian_hessian",
946 "linear_system_symbolic_factorization",
947 "fire_intermediate",
948 ] {
949 assert!(get(key).is_some(), "missing breakdown key {key}");
950 assert!(get(key).unwrap() >= 0.0, "negative time for {key}");
951 }
952
953 let total = get("linear_system_total").unwrap();
955 let sym = get("linear_system_symbolic_factorization").unwrap();
956 let fact = get("linear_system_factorization").unwrap();
957 let back = get("linear_system_back_solve").unwrap();
958 assert_eq!(total, sym + fact + back);
959 }
960
961 #[test]
968 fn report_carries_the_symbolic_and_intermediate_rows() {
969 let stats = TimingStatistics::new();
970 stats.linear_system_symbolic_factorization.start();
971 stats.linear_system_symbolic_factorization.end();
972 stats.fire_intermediate.start();
973 stats.fire_intermediate.end();
974
975 let text = stats.report();
976 assert!(
977 text.contains("LinearSystemSymbolicFactorization"),
978 "no symbolic-factorization row in:\n{text}"
979 );
980 assert!(
981 text.contains("FireIntermediateCallback"),
982 "no intermediate-callback row in:\n{text}"
983 );
984 }
985}