use std::marker::PhantomData;
use std::future::Future;
use std::sync::Arc;
use tokio::sync::{Mutex, RwLock};
use crate::{
resource::StorageConfig,
timing::{TimingTracker, time_async_fn},
utils::LassoError,
};
pub struct ADMMContext<G, S> {
pub config: ADMMConfig,
pub global: Arc<RwLock<G>>,
pub local: Arc<Mutex<S>>,
}
#[derive(Clone)]
pub struct ADMMConfig {
pub storage: StorageConfig,
}
pub trait ADMMProblem<G, S> {
fn context(&self) -> &ADMMContext<G, S>;
fn precompute(&self) -> impl Future<Output = Result<(), LassoError>>;
fn update_x(&self) -> impl Future<Output = Result<(), LassoError>>;
fn update_z(&self) -> impl Future<Output = Result<(), LassoError>>;
fn update_y(&self) -> impl Future<Output = Result<(), LassoError>>;
fn update_residuals(&self) -> impl Future<Output = Result<(), LassoError>>;
fn check_stopping_criteria(&self) -> impl Future<Output = Result<bool, LassoError>>;
}
pub struct ADMMSolver<G, S, P>
where
P: ADMMProblem<G, S>,
{
problem: P,
max_iter: usize,
timing_tracker: TimingTracker,
_global: PhantomData<G>,
_subproblem: PhantomData<S>,
}
impl<G, S, P> ADMMSolver<G, S, P>
where
P: ADMMProblem<G, S>,
{
pub fn new(problem: P, max_iter: usize) -> Self {
ADMMSolver {
problem,
max_iter,
timing_tracker: TimingTracker::new(),
_global: PhantomData,
_subproblem: PhantomData,
}
}
pub async fn solve(&mut self) -> Result<(), LassoError> {
time_async_fn(
&mut self.timing_tracker,
"precompute",
self.problem.precompute(),
)
.await?;
let mut i = 0;
loop {
self.timing_tracker.start_iteration();
println!("[ADMMSolver] ===== Iteration: {} =====", i);
time_async_fn(
&mut self.timing_tracker,
"update_x",
self.problem.update_x(),
)
.await?;
time_async_fn(
&mut self.timing_tracker,
"update_z",
self.problem.update_z(),
)
.await?;
time_async_fn(
&mut self.timing_tracker,
"update_y",
self.problem.update_y(),
)
.await?;
time_async_fn(
&mut self.timing_tracker,
"update_residuals",
self.problem.update_residuals(),
)
.await?;
let should_stop = time_async_fn(
&mut self.timing_tracker,
"check_stopping_criteria",
self.problem.check_stopping_criteria(),
)
.await?;
if should_stop {
break;
}
i += 1;
if i == self.max_iter {
break;
}
}
Ok(())
}
pub fn export_step_timings(&self, filename: &str) -> Result<(), LassoError> {
self.timing_tracker.write_step_timings_to_csv(filename)
}
pub fn export_lambda_timings(&self, filename: &str) -> Result<(), LassoError> {
self.timing_tracker.write_lambda_timings_to_csv(filename)
}
pub fn timing_tracker_mut(&mut self) -> &mut TimingTracker {
&mut self.timing_tracker
}
pub fn timing_tracker(&self) -> &TimingTracker {
&self.timing_tracker
}
pub fn export_all_timings(&self, filename_prefix: &str) -> Result<(), LassoError> {
let step_filename = format!("{}_steps.csv", filename_prefix);
let lambda_filename = format!("{}_lambdas.csv", filename_prefix);
self.export_step_timings(&step_filename)?;
self.export_lambda_timings(&lambda_filename)?;
println!("Exported step timings to: {}", step_filename);
println!("Exported lambda timings to: {}", lambda_filename);
Ok(())
}
pub fn print_timing_summary(&self) {
println!("\n=== ADMM Step Timing Summary ===");
let step_stats = self.timing_tracker.get_step_statistics();
for (step, (avg, max, count)) in step_stats {
println!(
"{}: avg={:.2}ms, max={:.2}ms, count={}",
step, avg, max, count
);
}
println!("\n=== Lambda Timing Summary ===");
let lambda_stats = self.timing_tracker.get_lambda_statistics();
for (lambda, (avg, max, count)) in lambda_stats {
println!(
"{}: avg={:.2}ms, max={:.2}ms, count={}",
lambda, avg, max, count
);
}
println!();
}
}