use crate::cost_curve::CostCurvePoint;
use crate::transfer::{ArmId, BetaParams, ContextBucket};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct BucketRegret {
pub best_mean: f32,
pub best_arm: ArmId,
pub cumulative_regret: f64,
pub observations: u64,
pub regret_history: Vec<f64>,
arm_means: HashMap<ArmId, (f64, u64)>,
}
impl BucketRegret {
fn new() -> Self {
Self {
best_mean: 0.0,
best_arm: ArmId("unknown".into()),
cumulative_regret: 0.0,
observations: 0,
regret_history: Vec::new(),
arm_means: HashMap::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegretTracker {
buckets: HashMap<ContextBucket, BucketRegret>,
pub total_regret: f64,
pub total_observations: u64,
snapshot_interval: u64,
}
impl RegretTracker {
pub fn new(snapshot_interval: u64) -> Self {
Self {
buckets: HashMap::new(),
total_regret: 0.0,
total_observations: 0,
snapshot_interval: snapshot_interval.max(1),
}
}
pub fn record(&mut self, bucket: &ContextBucket, arm: &ArmId, reward: f32) {
if !self.buckets.contains_key(bucket) {
self.buckets.insert(bucket.clone(), BucketRegret::new());
}
let entry = self.buckets.get_mut(bucket).unwrap();
if !entry.arm_means.contains_key(arm) {
entry.arm_means.insert(arm.clone(), (0.0, 0));
}
let (sum, count) = entry.arm_means.get_mut(arm).unwrap();
*sum += reward as f64;
*count += 1;
let arm_mean = *sum / *count as f64;
if arm_mean > entry.best_mean as f64 {
entry.best_mean = arm_mean as f32;
entry.best_arm = arm.clone();
}
let instant_regret = (entry.best_mean as f64 - reward as f64).max(0.0);
entry.cumulative_regret += instant_regret;
entry.observations += 1;
self.total_regret += instant_regret;
self.total_observations += 1;
if entry.observations % self.snapshot_interval == 0 {
entry.regret_history.push(entry.cumulative_regret);
}
}
pub fn regret_growth_rate(&self, bucket: &ContextBucket) -> Option<f32> {
let entry = self.buckets.get(bucket)?;
if entry.observations < 10 || entry.cumulative_regret < 1e-10 {
return None;
}
let log_regret = (entry.cumulative_regret).ln();
let log_t = (entry.observations as f64).ln();
Some((log_regret / log_t) as f32)
}
pub fn average_regret(&self) -> f64 {
if self.total_observations == 0 {
return 0.0;
}
self.total_regret / self.total_observations as f64
}
pub fn has_converged(&self, bucket: &ContextBucket, threshold: f32) -> bool {
self.regret_growth_rate(bucket)
.map_or(false, |rate| rate < threshold)
}
pub fn summary(&self) -> RegretSummary {
let bucket_rates: Vec<(ContextBucket, f32)> = self
.buckets
.keys()
.filter_map(|b| self.regret_growth_rate(b).map(|r| (b.clone(), r)))
.collect();
let mean_rate = if bucket_rates.is_empty() {
1.0
} else {
bucket_rates.iter().map(|(_, r)| r).sum::<f32>() / bucket_rates.len() as f32
};
RegretSummary {
total_regret: self.total_regret,
total_observations: self.total_observations,
average_regret: self.average_regret(),
mean_growth_rate: mean_rate,
bucket_count: self.buckets.len(),
converged_buckets: bucket_rates.iter().filter(|(_, r)| *r < 0.7).count(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RegretSummary {
pub total_regret: f64,
pub total_observations: u64,
pub average_regret: f64,
pub mean_growth_rate: f32,
pub bucket_count: usize,
pub converged_buckets: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DecayingBeta {
pub alpha: f32,
pub beta: f32,
pub decay_factor: f32,
pub effective_n: f32,
}
impl DecayingBeta {
pub fn new(decay_factor: f32) -> Self {
Self {
alpha: 1.0,
beta: 1.0,
decay_factor: decay_factor.clamp(0.9, 1.0),
effective_n: 0.0,
}
}
pub fn from_beta(params: &BetaParams, decay_factor: f32) -> Self {
Self {
alpha: params.alpha,
beta: params.beta,
decay_factor: decay_factor.clamp(0.9, 1.0),
effective_n: params.alpha + params.beta - 2.0,
}
}
pub fn update(&mut self, reward: f32) {
self.alpha = 1.0 + (self.alpha - 1.0) * self.decay_factor;
self.beta = 1.0 + (self.beta - 1.0) * self.decay_factor;
self.alpha += reward;
self.beta += 1.0 - reward;
self.effective_n = self.effective_n * self.decay_factor + 1.0;
}
pub fn mean(&self) -> f32 {
self.alpha / (self.alpha + self.beta)
}
pub fn variance(&self) -> f32 {
let total = self.alpha + self.beta;
(self.alpha * self.beta) / (total * total * (total + 1.0))
}
pub fn to_beta_params(&self) -> BetaParams {
BetaParams {
alpha: self.alpha,
beta: self.beta,
}
}
pub fn effective_window(&self) -> f32 {
if self.decay_factor >= 1.0 {
self.effective_n
} else {
1.0 / (1.0 - self.decay_factor)
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PlateauDetector {
pub window_size: usize,
pub improvement_threshold: f32,
pub consecutive_plateaus: u32,
pub total_plateaus: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub enum PlateauAction {
Continue,
IncreaseExploration,
TriggerTransfer,
InjectDiversity,
Reset,
}
impl PlateauDetector {
pub fn new(window_size: usize, improvement_threshold: f32) -> Self {
Self {
window_size: window_size.max(3),
improvement_threshold: improvement_threshold.max(0.001),
consecutive_plateaus: 0,
total_plateaus: 0,
}
}
pub fn check(&mut self, points: &[CostCurvePoint]) -> PlateauAction {
if points.len() < self.window_size * 2 {
self.consecutive_plateaus = 0;
return PlateauAction::Continue;
}
let n = points.len();
let recent = &points[n - self.window_size..];
let prior = &points[n - 2 * self.window_size..n - self.window_size];
let recent_mean = recent.iter().map(|p| p.accuracy).sum::<f32>() / recent.len() as f32;
let prior_mean = prior.iter().map(|p| p.accuracy).sum::<f32>() / prior.len() as f32;
let improvement = recent_mean - prior_mean;
if improvement.abs() < self.improvement_threshold {
self.consecutive_plateaus += 1;
self.total_plateaus += 1;
match self.consecutive_plateaus {
1 => PlateauAction::IncreaseExploration,
2..=3 => PlateauAction::TriggerTransfer,
4..=6 => PlateauAction::InjectDiversity,
_ => PlateauAction::Reset,
}
} else {
self.consecutive_plateaus = 0;
PlateauAction::Continue
}
}
pub fn check_cost(&self, points: &[CostCurvePoint]) -> bool {
if points.len() < self.window_size * 2 {
return false;
}
let n = points.len();
let recent = &points[n - self.window_size..];
let prior = &points[n - 2 * self.window_size..n - self.window_size];
let recent_cost =
recent.iter().map(|p| p.cost_per_solve).sum::<f32>() / recent.len() as f32;
let prior_cost = prior.iter().map(|p| p.cost_per_solve).sum::<f32>() / prior.len() as f32;
(prior_cost - recent_cost).abs() < self.improvement_threshold
}
pub fn learning_velocity(&self, points: &[CostCurvePoint]) -> f32 {
if points.len() < 2 {
return 0.0;
}
let n = points.len();
let window = self.window_size.min(n);
let recent = &points[n - window..];
if recent.len() < 2 {
return 0.0;
}
let first = recent.first().unwrap();
let last = recent.last().unwrap();
let dt = (last.cycle - first.cycle).max(1) as f32;
(last.accuracy - first.accuracy) / dt
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ParetoPoint {
pub kernel_id: String,
pub objectives: Vec<f32>,
pub generation: u32,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct ParetoFront {
pub front: Vec<ParetoPoint>,
pub evaluated: u64,
pub front_updates: u64,
}
impl ParetoFront {
pub fn new() -> Self {
Self::default()
}
pub fn dominates(a: &[f32], b: &[f32]) -> bool {
if a.len() != b.len() {
return false;
}
let mut at_least_equal = true;
let mut strictly_better = false;
for (ai, bi) in a.iter().zip(b.iter()) {
if ai < bi {
at_least_equal = false;
break;
}
if ai > bi {
strictly_better = true;
}
}
at_least_equal && strictly_better
}
pub fn insert(&mut self, point: ParetoPoint) -> bool {
self.evaluated += 1;
for existing in &self.front {
if Self::dominates(&existing.objectives, &point.objectives) {
return false; }
}
self.front
.retain(|existing| !Self::dominates(&point.objectives, &existing.objectives));
self.front.push(point);
self.front_updates += 1;
true
}
pub fn hypervolume(&self, reference: &[f32]) -> f32 {
if self.front.is_empty() || reference.is_empty() {
return 0.0;
}
let dim = reference.len();
if dim == 2 {
self.hypervolume_2d(reference)
} else {
self.front
.iter()
.map(|p| {
p.objectives
.iter()
.zip(reference.iter())
.map(|(oi, ri)| (oi - ri).max(0.0))
.product::<f32>()
})
.sum()
}
}
fn hypervolume_2d(&self, reference: &[f32]) -> f32 {
if self.front.is_empty() {
return 0.0;
}
let mut points: Vec<(f32, f32)> = self
.front
.iter()
.map(|p| {
let x = p.objectives.first().copied().unwrap_or(0.0);
let y = p.objectives.get(1).copied().unwrap_or(0.0);
(x, y)
})
.collect();
points.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal));
let ref_x = reference.first().copied().unwrap_or(0.0);
let ref_y = reference.get(1).copied().unwrap_or(0.0);
let mut volume = 0.0f32;
let mut prev_y = ref_y;
for &(x, y) in &points {
if y > prev_y {
volume += (x - ref_x) * (y - prev_y);
prev_y = y;
}
}
volume
}
pub fn len(&self) -> usize {
self.front.len()
}
pub fn is_empty(&self) -> bool {
self.front.is_empty()
}
pub fn best_on(&self, objective_index: usize) -> Option<&ParetoPoint> {
self.front.iter().max_by(|a, b| {
let va = a.objectives.get(objective_index).copied().unwrap_or(0.0);
let vb = b.objectives.get(objective_index).copied().unwrap_or(0.0);
va.partial_cmp(&vb).unwrap_or(std::cmp::Ordering::Equal)
})
}
pub fn spread(&self) -> Vec<f32> {
if self.front.is_empty() {
return Vec::new();
}
let dim = self.front[0].objectives.len();
(0..dim)
.map(|i| {
let vals: Vec<f32> = self.front.iter().map(|p| p.objectives[i]).collect();
let min = vals.iter().cloned().fold(f32::INFINITY, f32::min);
let max = vals.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
max - min
})
.collect()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CuriosityBonus {
visit_counts: HashMap<ContextBucket, HashMap<ArmId, u64>>,
pub total_visits: u64,
pub exploration_coeff: f32,
}
impl CuriosityBonus {
pub fn new(exploration_coeff: f32) -> Self {
Self {
visit_counts: HashMap::new(),
total_visits: 0,
exploration_coeff: exploration_coeff.max(0.0),
}
}
pub fn record_visit(&mut self, bucket: &ContextBucket, arm: &ArmId) {
if let Some(arms) = self.visit_counts.get_mut(bucket) {
if let Some(count) = arms.get_mut(arm) {
*count += 1;
} else {
arms.insert(arm.clone(), 1);
}
} else {
let mut arms = HashMap::new();
arms.insert(arm.clone(), 1u64);
self.visit_counts.insert(bucket.clone(), arms);
}
self.total_visits += 1;
}
pub fn bonus(&self, bucket: &ContextBucket, arm: &ArmId) -> f32 {
if self.total_visits < 2 {
return self.exploration_coeff; }
let arm_visits = self
.visit_counts
.get(bucket)
.and_then(|arms| arms.get(arm))
.copied()
.unwrap_or(0);
if arm_visits == 0 {
return self.exploration_coeff * 2.0; }
let log_n = (self.total_visits as f32).ln();
self.exploration_coeff * (log_n / arm_visits as f32).sqrt()
}
pub fn most_curious_bucket(&self) -> Option<&ContextBucket> {
let mut min_visits = u64::MAX;
let mut most_curious = None;
for (bucket, arms) in &self.visit_counts {
let total: u64 = arms.values().sum();
if total < min_visits {
min_visits = total;
most_curious = Some(bucket);
}
}
most_curious
}
pub fn novelty_score(&self, bucket: &ContextBucket) -> f32 {
if self.total_visits == 0 {
return 1.0;
}
let bucket_visits: u64 = self
.visit_counts
.get(bucket)
.map(|arms| arms.values().sum())
.unwrap_or(0);
if bucket_visits == 0 {
return 1.0;
}
1.0 - (bucket_visits as f32 / self.total_visits as f32)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MetaLearningEngine {
pub regret: RegretTracker,
pub plateau: PlateauDetector,
pub pareto: ParetoFront,
pub curiosity: CuriosityBonus,
pub decaying_betas: HashMap<(ContextBucket, ArmId), DecayingBeta>,
decay_factor: f32,
}
impl MetaLearningEngine {
pub fn new() -> Self {
Self {
regret: RegretTracker::new(50),
plateau: PlateauDetector::new(5, 0.005),
pareto: ParetoFront::new(),
curiosity: CuriosityBonus::new(1.41),
decaying_betas: HashMap::new(),
decay_factor: 0.995,
}
}
pub fn with_config(
regret_snapshot_interval: u64,
plateau_window: usize,
plateau_threshold: f32,
exploration_coeff: f32,
decay_factor: f32,
) -> Self {
Self {
regret: RegretTracker::new(regret_snapshot_interval),
plateau: PlateauDetector::new(plateau_window, plateau_threshold),
pareto: ParetoFront::new(),
curiosity: CuriosityBonus::new(exploration_coeff),
decaying_betas: HashMap::new(),
decay_factor,
}
}
pub fn record_decision(&mut self, bucket: &ContextBucket, arm: &ArmId, reward: f32) {
self.regret.record(bucket, arm, reward);
self.curiosity.record_visit(bucket, arm);
let key = (bucket.clone(), arm.clone());
if let Some(db) = self.decaying_betas.get_mut(&key) {
db.update(reward);
} else {
let mut db = DecayingBeta::new(self.decay_factor);
db.update(reward);
self.decaying_betas.insert(key, db);
}
}
pub fn record_kernel(
&mut self,
kernel_id: &str,
accuracy: f32,
cost: f32,
robustness: f32,
generation: u32,
) {
let point = ParetoPoint {
kernel_id: kernel_id.to_string(),
objectives: vec![accuracy, -cost, robustness],
generation,
};
self.pareto.insert(point);
}
pub fn check_plateau(&mut self, points: &[CostCurvePoint]) -> PlateauAction {
self.plateau.check(points)
}
pub fn boosted_score(&self, bucket: &ContextBucket, arm: &ArmId, thompson_sample: f32) -> f32 {
let bonus = self.curiosity.bonus(bucket, arm);
thompson_sample + bonus
}
pub fn decaying_mean(&self, bucket: &ContextBucket, arm: &ArmId) -> Option<f32> {
let key = (bucket.clone(), arm.clone());
self.decaying_betas.get(&key).map(|db| db.mean())
}
pub fn health_check(&self) -> MetaLearningHealth {
let regret_summary = self.regret.summary();
let pareto_size = self.pareto.len();
let is_learning = regret_summary.mean_growth_rate < 0.8;
let is_diverse = pareto_size >= 3;
let is_exploring = self.curiosity.total_visits > 0;
MetaLearningHealth {
regret: regret_summary,
pareto_size,
pareto_hypervolume: self.pareto.hypervolume(&[0.0, -1.0, 0.0]),
consecutive_plateaus: self.plateau.consecutive_plateaus,
total_plateaus: self.plateau.total_plateaus,
curiosity_total_visits: self.curiosity.total_visits,
decaying_beta_count: self.decaying_betas.len(),
is_learning,
is_diverse,
is_exploring,
}
}
}
impl Default for MetaLearningEngine {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MetaLearningHealth {
pub regret: RegretSummary,
pub pareto_size: usize,
pub pareto_hypervolume: f32,
pub consecutive_plateaus: u32,
pub total_plateaus: u32,
pub curiosity_total_visits: u64,
pub decaying_beta_count: usize,
pub is_learning: bool,
pub is_diverse: bool,
pub is_exploring: bool,
}
#[cfg(test)]
mod tests {
use super::*;
fn test_bucket(tier: &str, cat: &str) -> ContextBucket {
ContextBucket {
difficulty_tier: tier.into(),
category: cat.into(),
}
}
#[test]
fn test_regret_tracker_empty() {
let tracker = RegretTracker::new(10);
assert_eq!(tracker.total_regret, 0.0);
assert_eq!(tracker.average_regret(), 0.0);
}
#[test]
fn test_regret_tracker_optimal_arm() {
let mut tracker = RegretTracker::new(10);
let bucket = test_bucket("easy", "test");
let arm = ArmId("best".into());
for _ in 0..100 {
tracker.record(&bucket, &arm, 0.9);
}
assert_eq!(tracker.total_observations, 100);
assert!(tracker.total_regret < 1e-6);
}
#[test]
fn test_regret_tracker_suboptimal_arm() {
let mut tracker = RegretTracker::new(10);
let bucket = test_bucket("medium", "test");
let good = ArmId("good".into());
let bad = ArmId("bad".into());
for _ in 0..50 {
tracker.record(&bucket, &good, 0.9);
}
for _ in 0..50 {
tracker.record(&bucket, &bad, 0.3);
}
assert!(tracker.total_regret > 0.0);
assert!(tracker.average_regret() > 0.0);
}
#[test]
fn test_regret_growth_rate() {
let mut tracker = RegretTracker::new(5);
let bucket = test_bucket("hard", "test");
let arm_a = ArmId("a".into());
let arm_b = ArmId("b".into());
for _ in 0..50 {
tracker.record(&bucket, &arm_a, 0.8);
}
for _ in 0..50 {
tracker.record(&bucket, &arm_b, 0.4);
}
let rate = tracker.regret_growth_rate(&bucket);
assert!(rate.is_some());
}
#[test]
fn test_regret_summary() {
let mut tracker = RegretTracker::new(10);
let bucket = test_bucket("easy", "algo");
let arm = ArmId("test".into());
for _ in 0..20 {
tracker.record(&bucket, &arm, 0.7);
}
let summary = tracker.summary();
assert_eq!(summary.total_observations, 20);
assert_eq!(summary.bucket_count, 1);
}
#[test]
fn test_decaying_beta_initial() {
let db = DecayingBeta::new(0.995);
assert!((db.mean() - 0.5).abs() < 1e-6); assert_eq!(db.effective_n, 0.0);
}
#[test]
fn test_decaying_beta_update() {
let mut db = DecayingBeta::new(0.995);
for _ in 0..100 {
db.update(0.9); }
assert!(db.mean() > 0.7); assert!(db.effective_n > 50.0); }
#[test]
fn test_decaying_beta_adapts() {
let mut db = DecayingBeta::new(0.99);
for _ in 0..100 {
db.update(0.95);
}
let mean_after_good = db.mean();
assert!(mean_after_good > 0.8);
for _ in 0..100 {
db.update(0.1);
}
let mean_after_bad = db.mean();
assert!(mean_after_bad < mean_after_good);
assert!(mean_after_bad < 0.5); }
#[test]
fn test_decaying_beta_window() {
let db = DecayingBeta::new(0.99);
let window = db.effective_window();
assert!((window - 100.0).abs() < 1.0);
let db2 = DecayingBeta::new(0.995);
let window2 = db2.effective_window();
assert!((window2 - 200.0).abs() < 1.0); }
#[test]
fn test_decaying_to_standard() {
let mut db = DecayingBeta::new(0.995);
for _ in 0..10 {
db.update(0.8);
}
let params = db.to_beta_params();
assert!(params.alpha > 1.0);
assert!(params.beta > 1.0);
assert!((params.mean() - db.mean()).abs() < 1e-6);
}
#[test]
fn test_plateau_no_data() {
let mut detector = PlateauDetector::new(3, 0.01);
let action = detector.check(&[]);
assert_eq!(action, PlateauAction::Continue);
}
#[test]
fn test_plateau_not_enough_data() {
let mut detector = PlateauDetector::new(3, 0.01);
let points: Vec<CostCurvePoint> = (0..4)
.map(|i| CostCurvePoint {
cycle: i as u64,
accuracy: 0.5 + i as f32 * 0.1,
cost_per_solve: 0.1,
robustness: 0.8,
policy_violations: 0,
timestamp: i as f64,
})
.collect();
let action = detector.check(&points);
assert_eq!(action, PlateauAction::Continue);
}
#[test]
fn test_plateau_detected() {
let mut detector = PlateauDetector::new(3, 0.01);
let points: Vec<CostCurvePoint> = (0..6)
.map(|i| CostCurvePoint {
cycle: i as u64,
accuracy: 0.80 + (i as f32 * 0.001), cost_per_solve: 0.1,
robustness: 0.8,
policy_violations: 0,
timestamp: i as f64,
})
.collect();
let action = detector.check(&points);
assert_ne!(action, PlateauAction::Continue);
}
#[test]
fn test_plateau_improving() {
let mut detector = PlateauDetector::new(3, 0.01);
let points: Vec<CostCurvePoint> = (0..6)
.map(|i| CostCurvePoint {
cycle: i as u64,
accuracy: 0.50 + i as f32 * 0.08, cost_per_solve: 0.1,
robustness: 0.8,
policy_violations: 0,
timestamp: i as f64,
})
.collect();
let action = detector.check(&points);
assert_eq!(action, PlateauAction::Continue);
}
#[test]
fn test_plateau_escalation() {
let mut detector = PlateauDetector::new(3, 0.01);
let flat_points: Vec<CostCurvePoint> = (0..6)
.map(|i| CostCurvePoint {
cycle: i as u64,
accuracy: 0.80,
cost_per_solve: 0.1,
robustness: 0.8,
policy_violations: 0,
timestamp: i as f64,
})
.collect();
assert_eq!(
detector.check(&flat_points),
PlateauAction::IncreaseExploration
);
assert_eq!(detector.check(&flat_points), PlateauAction::TriggerTransfer);
assert_eq!(detector.check(&flat_points), PlateauAction::TriggerTransfer);
assert_eq!(detector.check(&flat_points), PlateauAction::InjectDiversity);
}
#[test]
fn test_learning_velocity() {
let detector = PlateauDetector::new(3, 0.01);
let points: Vec<CostCurvePoint> = (0..6)
.map(|i| CostCurvePoint {
cycle: i as u64,
accuracy: 0.50 + i as f32 * 0.1,
cost_per_solve: 0.1,
robustness: 0.8,
policy_violations: 0,
timestamp: i as f64,
})
.collect();
let velocity = detector.learning_velocity(&points);
assert!(velocity > 0.0); }
#[test]
fn test_pareto_dominates() {
assert!(ParetoFront::dominates(&[0.9, -0.1, 0.8], &[0.8, -0.2, 0.7]));
assert!(!ParetoFront::dominates(
&[0.9, -0.3, 0.8],
&[0.8, -0.1, 0.7]
));
assert!(!ParetoFront::dominates(
&[0.9, -0.1, 0.8],
&[0.9, -0.1, 0.8]
)); }
#[test]
fn test_pareto_insert_non_dominated() {
let mut front = ParetoFront::new();
assert!(front.insert(ParetoPoint {
kernel_id: "a".into(),
objectives: vec![0.9, -0.3, 0.7],
generation: 0,
}));
assert!(front.insert(ParetoPoint {
kernel_id: "b".into(),
objectives: vec![0.7, -0.1, 0.9],
generation: 0,
}));
assert_eq!(front.len(), 2);
}
#[test]
fn test_pareto_insert_dominated() {
let mut front = ParetoFront::new();
front.insert(ParetoPoint {
kernel_id: "good".into(),
objectives: vec![0.9, -0.1, 0.9],
generation: 0,
});
let added = front.insert(ParetoPoint {
kernel_id: "bad".into(),
objectives: vec![0.5, -0.5, 0.5],
generation: 0,
});
assert!(!added);
assert_eq!(front.len(), 1);
}
#[test]
fn test_pareto_removes_dominated() {
let mut front = ParetoFront::new();
front.insert(ParetoPoint {
kernel_id: "old".into(),
objectives: vec![0.5, -0.3, 0.5],
generation: 0,
});
front.insert(ParetoPoint {
kernel_id: "new".into(),
objectives: vec![0.9, -0.1, 0.9],
generation: 1,
});
assert_eq!(front.len(), 1);
assert_eq!(front.front[0].kernel_id, "new");
}
#[test]
fn test_pareto_best_on_objective() {
let mut front = ParetoFront::new();
front.insert(ParetoPoint {
kernel_id: "accurate".into(),
objectives: vec![0.95, -0.5, 0.6],
generation: 0,
});
front.insert(ParetoPoint {
kernel_id: "cheap".into(),
objectives: vec![0.7, -0.05, 0.7],
generation: 0,
});
front.insert(ParetoPoint {
kernel_id: "robust".into(),
objectives: vec![0.8, -0.3, 0.95],
generation: 0,
});
assert_eq!(front.best_on(0).unwrap().kernel_id, "accurate");
assert_eq!(front.best_on(1).unwrap().kernel_id, "cheap"); assert_eq!(front.best_on(2).unwrap().kernel_id, "robust");
}
#[test]
fn test_pareto_spread() {
let mut front = ParetoFront::new();
front.insert(ParetoPoint {
kernel_id: "a".into(),
objectives: vec![0.9, -0.5],
generation: 0,
});
front.insert(ParetoPoint {
kernel_id: "b".into(),
objectives: vec![0.5, -0.1],
generation: 0,
});
assert_eq!(front.len(), 2); let spread = front.spread();
assert_eq!(spread.len(), 2);
assert!((spread[0] - 0.4).abs() < 1e-4); assert!((spread[1] - 0.4).abs() < 1e-4); }
#[test]
fn test_pareto_hypervolume_2d() {
let mut front = ParetoFront::new();
front.insert(ParetoPoint {
kernel_id: "a".into(),
objectives: vec![1.0, 1.0],
generation: 0,
});
let hv = front.hypervolume(&[0.0, 0.0]);
assert!((hv - 1.0).abs() < 1e-4); }
#[test]
fn test_curiosity_bonus_unvisited() {
let curiosity = CuriosityBonus::new(1.41);
let bucket = test_bucket("hard", "novel");
let arm = ArmId("new".into());
let bonus = curiosity.bonus(&bucket, &arm);
assert!(bonus > 0.0); }
#[test]
fn test_curiosity_bonus_decays_with_visits() {
let mut curiosity = CuriosityBonus::new(1.41);
let bucket = test_bucket("easy", "test");
let arm = ArmId("a".into());
let bonus_before = curiosity.bonus(&bucket, &arm);
for _ in 0..50 {
curiosity.record_visit(&bucket, &arm);
}
let bonus_after = curiosity.bonus(&bucket, &arm);
assert!(bonus_after < bonus_before); }
#[test]
fn test_curiosity_novelty_score() {
let mut curiosity = CuriosityBonus::new(1.41);
let explored = test_bucket("easy", "common");
let novel = test_bucket("hard", "rare");
let arm = ArmId("a".into());
for _ in 0..100 {
curiosity.record_visit(&explored, &arm);
}
curiosity.record_visit(&novel, &arm);
let explored_novelty = curiosity.novelty_score(&explored);
let novel_novelty = curiosity.novelty_score(&novel);
assert!(novel_novelty > explored_novelty);
}
#[test]
fn test_meta_engine_creation() {
let engine = MetaLearningEngine::new();
assert_eq!(engine.regret.total_observations, 0);
assert!(engine.pareto.is_empty());
assert_eq!(engine.curiosity.total_visits, 0);
}
#[test]
fn test_meta_engine_record_decision() {
let mut engine = MetaLearningEngine::new();
let bucket = test_bucket("medium", "algo");
let arm = ArmId("greedy".into());
for _ in 0..50 {
engine.record_decision(&bucket, &arm, 0.85);
}
assert_eq!(engine.regret.total_observations, 50);
assert_eq!(engine.curiosity.total_visits, 50);
assert!(engine.decaying_mean(&bucket, &arm).unwrap() > 0.7);
}
#[test]
fn test_meta_engine_boosted_score() {
let mut engine = MetaLearningEngine::new();
let explored = test_bucket("easy", "common");
let novel = test_bucket("hard", "rare");
let arm = ArmId("a".into());
for _ in 0..100 {
engine.record_decision(&explored, &arm, 0.8);
}
let score_explored = engine.boosted_score(&explored, &arm, 0.5);
let score_novel = engine.boosted_score(&novel, &arm, 0.5);
assert!(score_novel > score_explored);
}
#[test]
fn test_meta_engine_kernel_recording() {
let mut engine = MetaLearningEngine::new();
engine.record_kernel("k1", 0.9, 0.3, 0.7, 0);
engine.record_kernel("k2", 0.7, 0.1, 0.9, 0);
engine.record_kernel("k3", 0.5, 0.5, 0.5, 0);
assert!(engine.pareto.len() <= 2);
}
#[test]
fn test_meta_engine_health_check() {
let mut engine = MetaLearningEngine::new();
let bucket = test_bucket("medium", "test");
let arm = ArmId("a".into());
for _ in 0..100 {
engine.record_decision(&bucket, &arm, 0.8);
}
let health = engine.health_check();
assert_eq!(health.curiosity_total_visits, 100);
assert!(health.is_exploring);
}
#[test]
fn test_meta_engine_plateau_check() {
let mut engine = MetaLearningEngine::new();
let flat_points: Vec<CostCurvePoint> = (0..10)
.map(|i| CostCurvePoint {
cycle: i as u64,
accuracy: 0.80,
cost_per_solve: 0.1,
robustness: 0.8,
policy_violations: 0,
timestamp: i as f64,
})
.collect();
let action = engine.check_plateau(&flat_points);
assert_ne!(action, PlateauAction::Continue);
}
#[test]
fn test_meta_engine_default() {
let engine = MetaLearningEngine::default();
assert_eq!(engine.curiosity.exploration_coeff, 1.41);
}
}