use crate::types::Number;
use crate::utils::{cpu_time, sys_time, wallclock_time};
use std::cell::Cell;
use std::rc::Rc;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DeadlineKind {
Wall,
Cpu,
}
#[derive(Debug, Clone)]
pub struct Deadline {
inner: Rc<DeadlineInner>,
}
#[derive(Debug)]
struct DeadlineInner {
wall_start: Number,
cpu_start: Number,
max_wall: Number,
max_cpu: Number,
}
impl Deadline {
pub fn new(max_wall: Number, max_cpu: Number) -> Self {
Self {
inner: Rc::new(DeadlineInner {
wall_start: wallclock_time(),
cpu_start: cpu_time(),
max_wall,
max_cpu,
}),
}
}
pub fn exceeded(&self) -> Option<DeadlineKind> {
if cpu_time() - self.inner.cpu_start >= self.inner.max_cpu {
return Some(DeadlineKind::Cpu);
}
if wallclock_time() - self.inner.wall_start >= self.inner.max_wall {
return Some(DeadlineKind::Wall);
}
None
}
pub fn max_wall(&self) -> Number {
self.inner.max_wall
}
pub fn max_cpu(&self) -> Number {
self.inner.max_cpu
}
pub fn remaining_wall(&self) -> Number {
self.inner.max_wall - (wallclock_time() - self.inner.wall_start)
}
pub fn remaining_cpu(&self) -> Number {
self.inner.max_cpu - (cpu_time() - self.inner.cpu_start)
}
}
#[derive(Debug)]
pub struct TimedTask {
enabled: Cell<bool>,
start_called: Cell<bool>,
end_called: Cell<bool>,
start_cpu: Cell<Number>,
start_sys: Cell<Number>,
start_wall: Cell<Number>,
total_cpu: Cell<Number>,
total_sys: Cell<Number>,
total_wall: Cell<Number>,
}
impl Default for TimedTask {
fn default() -> Self {
Self {
enabled: Cell::new(true),
start_called: Cell::new(false),
end_called: Cell::new(true),
start_cpu: Cell::new(0.0),
start_sys: Cell::new(0.0),
start_wall: Cell::new(0.0),
total_cpu: Cell::new(0.0),
total_sys: Cell::new(0.0),
total_wall: Cell::new(0.0),
}
}
}
impl TimedTask {
pub fn new() -> Self {
Self::default()
}
pub fn enable(&self) {
self.enabled.set(true);
}
pub fn disable(&self) {
self.enabled.set(false);
}
pub fn is_enabled(&self) -> bool {
self.enabled.get()
}
pub fn is_started(&self) -> bool {
self.start_called.get()
}
pub fn reset(&self) {
self.total_cpu.set(0.0);
self.total_sys.set(0.0);
self.total_wall.set(0.0);
self.start_called.set(false);
self.end_called.set(true);
}
pub fn start(&self) {
if !self.enabled.get() {
return;
}
self.end_called.set(false);
self.start_called.set(true);
self.start_cpu.set(cpu_time());
self.start_sys.set(sys_time());
self.start_wall.set(wallclock_time());
}
pub fn end(&self) {
if !self.enabled.get() {
return;
}
self.end_called.set(true);
self.start_called.set(false);
self.total_cpu
.set(self.total_cpu.get() + cpu_time() - self.start_cpu.get());
self.total_sys
.set(self.total_sys.get() + sys_time() - self.start_sys.get());
self.total_wall
.set(self.total_wall.get() + wallclock_time() - self.start_wall.get());
}
pub fn end_if_started(&self) {
if !self.enabled.get() {
return;
}
if self.start_called.get() {
self.end();
}
}
pub fn total_cpu_time(&self) -> Number {
self.total_cpu.get()
}
pub fn total_sys_time(&self) -> Number {
self.total_sys.get()
}
pub fn total_wallclock_time(&self) -> Number {
self.total_wall.get()
}
pub fn live_wallclock_time(&self) -> Number {
if self.enabled.get() && self.start_called.get() {
self.total_wall.get() + wallclock_time() - self.start_wall.get()
} else {
self.total_wall.get()
}
}
pub fn live_cpu_time(&self) -> Number {
if self.enabled.get() && self.start_called.get() {
self.total_cpu.get() + cpu_time() - self.start_cpu.get()
} else {
self.total_cpu.get()
}
}
pub fn guard(&self) -> TimedGuard<'_> {
self.start();
TimedGuard { task: Some(self) }
}
}
#[must_use = "the guard ends the timer when dropped; bind it to a variable"]
pub struct TimedGuard<'a> {
task: Option<&'a TimedTask>,
}
impl<'a> TimedGuard<'a> {
pub fn stop(mut self) {
if let Some(t) = self.task.take() {
t.end_if_started();
}
}
}
impl<'a> Drop for TimedGuard<'a> {
fn drop(&mut self) {
if let Some(t) = self.task.take() {
t.end_if_started();
}
}
}
#[derive(Debug, Default)]
pub struct TimingStatistics {
pub overall_alg: TimedTask,
pub print_problem_statistics: TimedTask,
pub initialize_iterates: TimedTask,
pub update_hessian: TimedTask,
pub output_iteration: TimedTask,
pub update_barrier_parameter: TimedTask,
pub compute_search_direction: TimedTask,
pub compute_acceptable_trial_point: TimedTask,
pub accept_trial_point: TimedTask,
pub check_convergence: TimedTask,
pub linear_system_factorization: TimedTask,
pub linear_system_back_solve: TimedTask,
pub linear_system_structure_converter: TimedTask,
pub linear_system_structure_converter_init: TimedTask,
pub quality_function_search: TimedTask,
pub total_callback_time: TimedTask,
pub total_function_evaluation_time: TimedTask,
pub eval_obj: TimedTask,
pub eval_grad_obj: TimedTask,
pub eval_constr: TimedTask,
pub eval_constr_jac: TimedTask,
pub eval_lag_hess: TimedTask,
}
impl TimingStatistics {
pub fn new() -> Self {
Self::default()
}
pub fn report(&self) -> String {
use std::fmt::Write as _;
let mut s = String::new();
let row = |s: &mut String, label: &str, t: &TimedTask| {
let _ = writeln!(
s,
"{label:<42} {wall:>10.3}s",
wall = t.total_wallclock_time()
);
};
s.push_str("\nTiming Statistics:\n");
row(
&mut s,
"OverallAlgorithm....................:",
&self.overall_alg,
);
row(
&mut s,
" InitializeIterates.................:",
&self.initialize_iterates,
);
row(
&mut s,
" UpdateHessian......................:",
&self.update_hessian,
);
row(
&mut s,
" OutputIteration....................:",
&self.output_iteration,
);
row(
&mut s,
" UpdateBarrierParameter.............:",
&self.update_barrier_parameter,
);
row(
&mut s,
" ComputeSearchDirection.............:",
&self.compute_search_direction,
);
row(
&mut s,
" ComputeAcceptableTrialPoint........:",
&self.compute_acceptable_trial_point,
);
row(
&mut s,
" AcceptTrialPoint...................:",
&self.accept_trial_point,
);
row(
&mut s,
" CheckConvergence...................:",
&self.check_convergence,
);
row(
&mut s,
"LinearSystemFactorization...........:",
&self.linear_system_factorization,
);
row(
&mut s,
"LinearSystemBackSolve...............:",
&self.linear_system_back_solve,
);
row(
&mut s,
"QualityFunctionSearch...............:",
&self.quality_function_search,
);
row(
&mut s,
"TotalFunctionEvaluations............:",
&self.total_function_evaluation_time,
);
row(
&mut s,
" ObjectiveFunctionEvaluations.......:",
&self.eval_obj,
);
row(
&mut s,
" ObjectiveGradientEvaluations.......:",
&self.eval_grad_obj,
);
row(
&mut s,
" ConstraintEvaluations..............:",
&self.eval_constr,
);
row(
&mut s,
" ConstraintJacobianEvaluations......:",
&self.eval_constr_jac,
);
row(
&mut s,
" LagrangianHessianEvaluations.......:",
&self.eval_lag_hess,
);
s
}
pub fn wall_time_breakdown(&self) -> Vec<(&'static str, Number)> {
let factorization = self.linear_system_factorization.total_wallclock_time();
let back_solve = self.linear_system_back_solve.total_wallclock_time();
vec![
("overall_alg", self.overall_alg.total_wallclock_time()),
("update_hessian", self.update_hessian.total_wallclock_time()),
(
"compute_search_direction",
self.compute_search_direction.total_wallclock_time(),
),
("linear_system_total", factorization + back_solve),
("linear_system_factorization", factorization),
("linear_system_back_solve", back_solve),
(
"function_evaluations_total",
self.total_function_evaluation_time.total_wallclock_time(),
),
("eval_objective", self.eval_obj.total_wallclock_time()),
("eval_gradient", self.eval_grad_obj.total_wallclock_time()),
("eval_constraints", self.eval_constr.total_wallclock_time()),
(
"eval_constraint_jacobian",
self.eval_constr_jac.total_wallclock_time(),
),
(
"eval_lagrangian_hessian",
self.eval_lag_hess.total_wallclock_time(),
),
(
"total_callback",
self.total_callback_time.total_wallclock_time(),
),
]
}
pub fn set_detailed_enabled(&self, on: bool) {
let set = |t: &TimedTask| {
if on {
t.enable();
} else {
t.disable();
}
};
set(&self.print_problem_statistics);
set(&self.initialize_iterates);
set(&self.update_hessian);
set(&self.output_iteration);
set(&self.update_barrier_parameter);
set(&self.compute_search_direction);
set(&self.compute_acceptable_trial_point);
set(&self.accept_trial_point);
set(&self.check_convergence);
set(&self.linear_system_factorization);
set(&self.linear_system_back_solve);
set(&self.linear_system_structure_converter);
set(&self.linear_system_structure_converter_init);
set(&self.quality_function_search);
set(&self.total_callback_time);
set(&self.total_function_evaluation_time);
set(&self.eval_obj);
set(&self.eval_grad_obj);
set(&self.eval_constr);
set(&self.eval_constr_jac);
set(&self.eval_lag_hess);
}
pub fn reset(&self) {
self.overall_alg.reset();
self.print_problem_statistics.reset();
self.initialize_iterates.reset();
self.update_hessian.reset();
self.output_iteration.reset();
self.update_barrier_parameter.reset();
self.compute_search_direction.reset();
self.compute_acceptable_trial_point.reset();
self.accept_trial_point.reset();
self.check_convergence.reset();
self.linear_system_factorization.reset();
self.linear_system_back_solve.reset();
self.linear_system_structure_converter.reset();
self.linear_system_structure_converter_init.reset();
self.quality_function_search.reset();
self.total_callback_time.reset();
self.total_function_evaluation_time.reset();
self.eval_obj.reset();
self.eval_grad_obj.reset();
self.eval_constr.reset();
self.eval_constr_jac.reset();
self.eval_lag_hess.reset();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deadline_unbounded_never_trips() {
let d = Deadline::new(1e6, 1e6);
assert!(d.exceeded().is_none());
assert_eq!(d.max_wall(), 1e6);
assert_eq!(d.max_cpu(), 1e6);
}
#[test]
fn deadline_zero_wall_trips_wall() {
let d = Deadline::new(0.0, 1e6);
for _ in 0..10_000 {
if d.exceeded().is_some() {
break;
}
std::hint::black_box(0u64);
}
assert_eq!(d.exceeded(), Some(DeadlineKind::Wall));
}
#[test]
fn deadline_zero_cpu_takes_priority() {
let d = Deadline::new(0.0, 0.0);
for _ in 0..10_000 {
if d.exceeded().is_some() {
break;
}
std::hint::black_box(0u64);
}
assert_eq!(d.exceeded(), Some(DeadlineKind::Cpu));
}
#[test]
fn deadline_remaining_reports_budget_left() {
let d = Deadline::new(1e6, 1e6);
let rw = d.remaining_wall();
let rc = d.remaining_cpu();
assert!(rw > 0.0 && rw <= 1e6, "remaining_wall out of range: {rw}");
assert!(rc > 0.0 && rc <= 1e6, "remaining_cpu out of range: {rc}");
}
#[test]
fn deadline_remaining_goes_nonpositive_once_crossed() {
let d = Deadline::new(0.0, 1e6);
for _ in 0..10_000 {
if d.exceeded().is_some() {
break;
}
std::hint::black_box(0u64);
}
assert!(
d.remaining_wall() <= 0.0,
"remaining_wall should be non-positive once the budget is crossed"
);
}
#[test]
fn deadline_is_cheaply_clonable_and_shares_start() {
let d = Deadline::new(1e6, 1e6);
let d2 = d.clone();
assert_eq!(d2.max_wall(), d.max_wall());
assert_eq!(d2.max_cpu(), d.max_cpu());
assert!(d2.exceeded().is_none());
}
#[test]
fn start_end_accumulates_nonneg() {
let t = TimedTask::new();
t.start();
for _ in 0..1000 {
std::hint::black_box(0u64);
}
t.end();
assert!(t.total_wallclock_time() >= 0.0);
}
#[test]
fn disabled_is_noop() {
let t = TimedTask::new();
t.disable();
t.start();
t.end();
assert_eq!(t.total_wallclock_time(), 0.0);
}
#[test]
fn set_detailed_enabled_gates_all_but_overall_alg() {
let stats = TimingStatistics::new();
assert!(stats.overall_alg.is_enabled());
assert!(stats.eval_obj.is_enabled());
assert!(stats.check_convergence.is_enabled());
stats.set_detailed_enabled(false);
assert!(stats.overall_alg.is_enabled(), "overall_alg must stay live");
assert!(!stats.eval_obj.is_enabled());
assert!(!stats.check_convergence.is_enabled());
assert!(!stats.total_function_evaluation_time.is_enabled());
assert!(!stats.linear_system_factorization.is_enabled());
stats.eval_obj.start();
stats.eval_obj.end();
assert_eq!(stats.eval_obj.total_wallclock_time(), 0.0);
stats.set_detailed_enabled(true);
assert!(stats.eval_obj.is_enabled());
assert!(stats.check_convergence.is_enabled());
}
#[test]
fn end_if_started_handles_unstarted() {
let t = TimedTask::new();
t.end_if_started();
assert_eq!(t.total_wallclock_time(), 0.0);
}
#[test]
fn wall_time_breakdown_reports_subsystems() {
let stats = TimingStatistics::new();
stats.linear_system_factorization.start();
stats.linear_system_factorization.end();
stats.eval_lag_hess.start();
stats.eval_lag_hess.end();
let bd = stats.wall_time_breakdown();
let get = |k: &str| bd.iter().find(|(label, _)| *label == k).map(|(_, v)| *v);
for key in [
"overall_alg",
"linear_system_total",
"linear_system_factorization",
"linear_system_back_solve",
"function_evaluations_total",
"eval_objective",
"eval_gradient",
"eval_constraints",
"eval_constraint_jacobian",
"eval_lagrangian_hessian",
] {
assert!(get(key).is_some(), "missing breakdown key {key}");
assert!(get(key).unwrap() >= 0.0, "negative time for {key}");
}
let total = get("linear_system_total").unwrap();
let fact = get("linear_system_factorization").unwrap();
let back = get("linear_system_back_solve").unwrap();
assert_eq!(total, fact + back);
}
}