#[cfg(feature = "counters")]
use crate::glmm::PIRLS_MAX_ITERS;
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub enum Stage {
One = 0,
Two = 1,
}
#[cfg(feature = "counters")]
pub const PIRLS_HIST_LEN: usize = PIRLS_MAX_ITERS + 1;
#[cfg(feature = "counters")]
#[derive(Clone, Copy, Debug)]
pub struct EvalCounters {
pub stage_evals: [u32; 2],
pub stage_last_improve: [u32; 2],
stage_best: [f64; 2],
pub pirls_hist: [u32; PIRLS_HIST_LEN],
pending_pirls_iters: u32,
pub agq_evals: u32,
pub agq_node_evals: u64,
pub nb_nodes: u32,
pub nb_evals_total: u64,
}
#[cfg(feature = "counters")]
impl EvalCounters {
pub(crate) fn new() -> Self {
EvalCounters {
stage_evals: [0; 2],
stage_last_improve: [0; 2],
stage_best: [f64::INFINITY; 2],
pirls_hist: [0; PIRLS_HIST_LEN],
pending_pirls_iters: 0,
agq_evals: 0,
agq_node_evals: 0,
nb_nodes: 0,
nb_evals_total: 0,
}
}
pub(crate) fn reset(&mut self) {
*self = Self::new();
}
pub(crate) fn record_eval(&mut self, stage: Stage, obj: f64) {
let s = stage as usize;
self.stage_evals[s] += 1;
if obj < self.stage_best[s] {
self.stage_best[s] = obj;
self.stage_last_improve[s] = self.stage_evals[s];
}
}
pub(crate) fn set_pirls_iters(&mut self, iters: usize) {
self.pending_pirls_iters = iters as u32;
}
pub(crate) fn commit_pirls_iters(&mut self) {
let bucket = (self.pending_pirls_iters as usize).min(PIRLS_MAX_ITERS);
self.pirls_hist[bucket] += 1;
self.pending_pirls_iters = 0;
}
pub(crate) fn record_agq_eval(&mut self, nodes: u64) {
self.agq_evals += 1;
self.agq_node_evals += nodes;
}
pub(crate) fn record_nb_node(&mut self, n_eval: usize) {
self.nb_nodes += 1;
self.nb_evals_total += n_eval as u64;
}
pub fn evals_after_last_improve(&self, stage: Stage) -> u32 {
let s = stage as usize;
self.stage_evals[s] - self.stage_last_improve[s]
}
}
#[cfg(not(feature = "counters"))]
#[derive(Clone, Copy, Debug)]
pub struct EvalCounters;
#[cfg(not(feature = "counters"))]
impl EvalCounters {
#[inline(always)]
pub(crate) fn new() -> Self {
EvalCounters
}
#[inline(always)]
pub(crate) fn reset(&mut self) {}
#[inline(always)]
pub(crate) fn record_eval(&mut self, _stage: Stage, _obj: f64) {}
#[inline(always)]
pub(crate) fn set_pirls_iters(&mut self, _iters: usize) {}
#[inline(always)]
pub(crate) fn commit_pirls_iters(&mut self) {}
#[inline(always)]
pub(crate) fn record_agq_eval(&mut self, _nodes: u64) {}
#[inline(always)]
pub(crate) fn record_nb_node(&mut self, _n_eval: usize) {}
}