use ndarray::Array1;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Mutex, OnceLock};
#[derive(Clone, Debug)]
pub struct OuterEvalRecord {
pub theta: Array1<f64>,
pub cost: f64,
pub gradient: Array1<f64>,
}
const MAX_CAPTURED: usize = 8;
static ENABLED: AtomicBool = AtomicBool::new(false);
fn buffer() -> &'static Mutex<Vec<OuterEvalRecord>> {
static BUFFER: OnceLock<Mutex<Vec<OuterEvalRecord>>> = OnceLock::new();
BUFFER.get_or_init(|| Mutex::new(Vec::new()))
}
pub fn enable_outer_eval_capture() {
buffer().lock().expect("outer-eval capture buffer").clear();
ENABLED.store(true, Ordering::Relaxed);
}
pub fn take_outer_eval_capture() -> Vec<OuterEvalRecord> {
ENABLED.store(false, Ordering::Relaxed);
std::mem::take(&mut *buffer().lock().expect("outer-eval capture buffer"))
}
pub(crate) fn record_outer_eval(theta: &Array1<f64>, cost: f64, gradient: &Array1<f64>) {
if !ENABLED.load(Ordering::Relaxed) {
return;
}
let mut b = buffer().lock().expect("outer-eval capture buffer");
if b.len() < MAX_CAPTURED {
b.push(OuterEvalRecord {
theta: theta.clone(),
cost,
gradient: gradient.clone(),
});
}
}