use std::time::Duration;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub enum StopReason {
TargetVocabReached,
MaxIterationsReached,
NoEligiblePairs,
PlateauReached,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct IterationMetrics {
pub iteration: usize,
pub best_frequency: usize,
pub merges_applied: usize,
pub distinct_pairs: usize,
pub elapsed_iteration: Duration,
pub elapsed_total: Duration,
pub rss_kb: Option<usize>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct TrainingMetrics {
pub iterations: Vec<IterationMetrics>,
pub total_duration: Duration,
pub stop_reason: StopReason,
}
impl TrainingMetrics {
#[must_use]
pub fn new(capacity: usize) -> Self {
Self {
iterations: Vec::with_capacity(capacity),
total_duration: Duration::ZERO,
stop_reason: StopReason::TargetVocabReached,
}
}
}
#[cfg(target_os = "linux")]
fn current_rss_kb() -> Option<usize> {
use std::fs::File;
use std::io::{BufRead, BufReader};
let file = File::open("/proc/self/status").ok()?;
for line in BufReader::new(file).lines().map_while(Result::ok) {
if let Some(rest) = line.strip_prefix("VmRSS:") {
let value = rest
.split_whitespace()
.find_map(|part| part.parse::<usize>().ok());
return value;
}
}
None
}
#[cfg(not(target_os = "linux"))]
fn current_rss_kb() -> Option<usize> {
None
}
pub fn sample_rss_kb() -> Option<usize> {
current_rss_kb()
}