use scirs2_core::ndarray::Array1;
use scirs2_core::numeric::Float;
use std::collections::{HashMap, VecDeque};
use std::time::{Duration, Instant};
use crate::error::Result;
use crate::utils::scalar_or;
#[derive(Debug, Clone)]
pub enum PerformanceMetric<A: Float + Send + Sync> {
Loss(A),
Accuracy(A),
F1Score(A),
AUC(A),
Custom { name: String, value: A },
}
#[derive(Debug, Clone)]
pub struct EnhancedAdaptiveLRController<A: Float + Send + Sync> {
current_lr: A,
base_lr: A,
adaptation_strategy: MultiSignalAdaptationStrategy<A>,
gradient_adapter: GradientBasedAdapter<A>,
performance_adapter: PerformanceBasedAdapter<A>,
drift_adapter: DriftAwareAdapter<A>,
resource_adapter: ResourceAwareAdapter<A>,
meta_optimizer: MetaOptimizer<A>,
adaptation_history: VecDeque<AdaptationEvent<A>>,
config: AdaptiveLRConfig<A>,
}
#[derive(Debug, Clone)]
pub struct AdaptiveLRConfig<A: Float + Send + Sync> {
pub base_lr: A,
pub min_lr: A,
pub max_lr: A,
pub enable_gradient_adaptation: bool,
pub enable_performance_adaptation: bool,
pub enable_drift_adaptation: bool,
pub enable_resource_adaptation: bool,
pub enable_meta_learning: bool,
pub history_window_size: usize,
pub adaptation_frequency: usize,
pub adaptation_sensitivity: A,
pub use_ensemble_voting: bool,
pub step_time_budget: Option<Duration>,
pub memory_budget_mb: Option<f64>,
}
#[derive(Debug, Clone)]
pub struct MultiSignalAdaptationStrategy<A: Float + Send + Sync> {
pub(crate) signal_weights: HashMap<AdaptationSignalType, A>,
pub(crate) voting_history: VecDeque<SignalVote<A>>,
pub(crate) conflict_resolution: ConflictResolution,
pub(crate) signal_reliability: HashMap<AdaptationSignalType, A>,
pub(crate) last_decision: Option<AdaptationDecision<A>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum AdaptationSignalType {
GradientMagnitude,
GradientVariance,
LossProgression,
AccuracyTrend,
ConceptDrift,
ResourceUtilization,
ModelComplexity,
DataQuality,
}
#[derive(Debug, Clone)]
pub struct SignalVote<A: Float + Send + Sync> {
signal_type: AdaptationSignalType,
recommended_lr_change: A, confidence: A,
reasoning: String,
timestamp: Instant,
}
impl<A: Float + Send + Sync> SignalVote<A> {
pub fn signal_type(&self) -> AdaptationSignalType {
self.signal_type
}
pub fn recommended_lr_change(&self) -> A {
self.recommended_lr_change
}
pub fn confidence(&self) -> A {
self.confidence
}
pub fn reasoning(&self) -> &str {
&self.reasoning
}
pub fn timestamp(&self) -> Instant {
self.timestamp
}
}
#[derive(Debug, Clone, Copy)]
pub enum ConflictResolution {
WeightedAverage,
HighestConfidence,
MajorityVote { threshold: f64 },
Conservative,
MetaLearned,
}
#[derive(Debug, Clone)]
pub struct AdaptationDecision<A: Float + Send + Sync> {
new_lr: A,
lr_multiplier: A,
contributing_signals: Vec<AdaptationSignalType>,
confidence: A,
rationale: String,
timestamp: Instant,
}
impl<A: Float + Send + Sync> AdaptationDecision<A> {
pub fn new_lr(&self) -> A {
self.new_lr
}
pub fn lr_multiplier(&self) -> A {
self.lr_multiplier
}
pub fn contributing_signals(&self) -> &[AdaptationSignalType] {
&self.contributing_signals
}
pub fn confidence(&self) -> A {
self.confidence
}
pub fn rationale(&self) -> &str {
&self.rationale
}
pub fn timestamp(&self) -> Instant {
self.timestamp
}
}
#[derive(Debug, Clone)]
pub struct GradientBasedAdapter<A: Float + Send + Sync> {
magnitude_history: VecDeque<A>,
direction_variance_history: VecDeque<A>,
norm_statistics: GradientNormStatistics<A>,
snr_estimator: SignalToNoiseEstimator<A>,
staleness_detector: GradientStalenessDetector,
}
#[derive(Debug, Clone)]
pub struct PerformanceBasedAdapter<A: Float + Send + Sync> {
metric_history: HashMap<String, VecDeque<A>>,
trend_analyzer: PerformanceTrendAnalyzer<A>,
plateau_detector: PlateauDetector<A>,
overfitting_detector: OverfittingDetector<A>,
efficiency_tracker: LearningEfficiencyTracker<A>,
}
#[derive(Debug, Clone)]
pub struct DriftAwareAdapter<A: Float + Send + Sync> {
drift_detectors: Vec<ConceptDriftDetector<A>>,
distribution_tracker: DistributionTracker<A>,
adaptation_speed: AdaptationSpeedController<A>,
drift_severity: DriftSeverityAssessor<A>,
}
#[derive(Debug, Clone)]
pub struct ResourceAwareAdapter<A: Float + Send + Sync> {
memory_tracker: MemoryUsageTracker,
compute_tracker: ComputationTimeTracker,
energy_tracker: EnergyConsumptionTracker,
throughput_requirements: ThroughputRequirements<A>,
budget_manager: ResourceBudgetManager<A>,
}
#[derive(Debug, Clone)]
pub struct MetaOptimizer<A: Float + Send + Sync> {
optimization_history: VecDeque<HyperparameterUpdate<A>>,
exploration_strategy: ExplorationStrategy<A>,
transfer_learner: TransferLearner<A>,
}
#[derive(Debug, Clone)]
pub struct AdaptationEvent<A: Float + Send + Sync> {
timestamp: Instant,
old_lr: A,
new_lr: A,
trigger_signals: Vec<AdaptationSignalType>,
effectiveness_score: Option<A>, }
#[derive(Debug, Clone)]
pub struct GradientNormStatistics<A: Float + Send + Sync> {
mean: A,
variance: A,
skewness: A,
kurtosis: A,
percentiles: Vec<A>, autocorrelation: A,
}
#[derive(Debug, Clone)]
pub struct SignalToNoiseEstimator<A: Float + Send + Sync> {
signal_estimate: A,
noise_estimate: A,
snr_history: VecDeque<A>,
}
#[derive(Debug, Clone, Copy)]
pub enum SNREstimationMethod {
MovingAverage,
ExponentialSmoothing,
RobustEstimation,
WaveletDenoising,
}
#[derive(Debug, Clone, Default)]
pub struct GradientStalenessDetector {
gradient_timestamps: VecDeque<Instant>,
}
#[derive(Debug, Clone)]
pub struct PerformanceTrendAnalyzer<A: Float + Send + Sync> {
trend_detection_window: usize,
trend_types: Vec<TrendType>,
trend_strength: A,
}
#[derive(Debug, Clone, Copy)]
pub enum TrendType {
Improving,
Degrading,
Oscillating,
Plateau,
Volatile,
}
#[derive(Debug, Clone)]
pub struct PlateauDetector<A: Float + Send + Sync> {
plateau_threshold: A,
min_plateau_duration: usize,
current_plateau_length: usize,
plateau_confidence: A,
}
#[derive(Debug, Clone)]
pub struct OverfittingDetector<A: Float + Send + Sync> {
train_loss_history: VecDeque<A>,
val_loss_history: VecDeque<A>,
}
#[derive(Debug, Clone)]
pub struct LearningEfficiencyTracker<A: Float + Send + Sync> {
loss_reduction_per_step: VecDeque<A>,
efficiency_score: A,
efficiency_trend: TrendType,
}
#[derive(Debug, Clone)]
pub struct ConceptDriftDetector<A: Float + Send + Sync> {
pub(crate) detection_method: DriftDetectionMethod,
pub(crate) drift_threshold: A,
pub(crate) window_size: usize,
pub(crate) drift_confidence: A,
pub(crate) last_drift_time: Option<Instant>,
pub(crate) inner: LossDriftDetector<A>,
}
#[derive(Debug, Clone, Copy)]
pub enum DriftDetectionMethod {
ADWIN,
DDM,
EDDM,
PageHinkley,
KSWIN,
Statistical,
}
#[derive(Debug, Clone)]
pub struct DistributionTracker<A: Float + Send + Sync> {
feature_distributions: HashMap<usize, FeatureDistribution<A>>,
distribution_drift_score: A,
}
#[derive(Debug, Clone)]
pub struct FeatureDistribution<A: Float + Send + Sync> {
mean: A,
variance: A,
histogram: Vec<A>,
last_update: Instant,
}
#[derive(Debug, Clone)]
pub struct AdaptationSpeedController<A: Float + Send + Sync> {
base_adaptation_rate: A,
current_adaptation_rate: A,
acceleration_factor: A,
deceleration_factor: A,
momentum: A,
}
#[derive(Debug, Clone)]
pub struct DriftSeverityAssessor<A: Float + Send + Sync> {
severity_levels: Vec<DriftSeverityLevel<A>>,
current_severity: DriftSeverityLevel<A>,
severity_history: VecDeque<DriftSeverityLevel<A>>,
}
#[derive(Debug, Clone)]
pub struct DriftSeverityLevel<A: Float + Send + Sync> {
level: DriftSeverity,
recommended_lr_adjustment: A,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum DriftSeverity {
None,
Mild,
Moderate,
Severe,
Critical,
}
#[derive(Debug, Clone, Default)]
pub struct MemoryUsageTracker {
pub(crate) current_usage_mb: f64,
pub(crate) peak_usage_mb: f64,
pub(crate) usage_history: VecDeque<f64>,
pub(crate) memory_pressure: Option<f64>,
}
#[derive(Debug, Clone, Default)]
pub struct ComputationTimeTracker {
pub(crate) step_times: VecDeque<Duration>,
pub(crate) average_step_time: Duration,
pub(crate) time_budget: Option<Duration>,
pub(crate) time_pressure: Option<f64>,
}
#[derive(Debug, Clone, Default)]
pub struct EnergyConsumptionTracker {
pub(crate) energy_per_step: VecDeque<f64>,
pub(crate) cumulative_energy: f64,
pub(crate) energy_efficiency: Option<f64>,
}
#[derive(Debug, Clone)]
pub struct ThroughputRequirements<A: Float + Send + Sync> {
min_samples_per_second: A,
current_throughput: A,
throughput_deficit: A,
}
#[derive(Debug, Clone)]
pub struct ResourceBudgetManager<A: Float + Send + Sync> {
memory_budget_mb: f64,
compute_budget_seconds: f64,
budget_utilization: A,
budget_violations: usize,
}
#[derive(Debug, Clone)]
pub struct HyperparameterUpdate<A: Float + Send + Sync> {
features: Array1<A>,
reward: A, }
#[derive(Debug, Clone)]
pub struct ExplorationStrategy<A: Float + Send + Sync> {
exploration_rate: A,
arm_rewards: HashMap<usize, A>,
arm_counts: HashMap<usize, usize>,
}
#[derive(Debug, Clone, Copy)]
pub enum ExplorationStrategyType {
EpsilonGreedy,
UCB1,
ThompsonSampling,
LinUCB,
ContextualBandit,
}
#[derive(Debug, Clone)]
pub struct TransferLearner<A: Float + Send + Sync> {
source_task_data: Vec<TaskData<A>>,
transfer_confidence: A,
}
#[derive(Debug, Clone)]
pub struct TaskData<A: Float + Send + Sync> {
optimal_lr_sequence: Vec<A>,
}
#[derive(Debug, Clone, Default)]
pub struct AdaptationStatistics<A: Float + Send + Sync> {
pub total_adaptations: usize,
pub successful_adaptations: usize,
pub avg_adaptation_frequency: A,
pub lr_volatility: A,
pub signal_reliability_scores: HashMap<AdaptationSignalType, A>,
pub signal_effectiveness: HashMap<AdaptationSignalType, A>,
pub resource_efficiency_gains: A,
pub convergence_speed_improvement: A,
}
impl<A: Float + Default + Clone + std::iter::Sum + Send + Sync> EnhancedAdaptiveLRController<A> {
pub fn new(config: AdaptiveLRConfig<A>) -> Result<Self> {
let adaptation_strategy = MultiSignalAdaptationStrategy::new(&config)?;
let gradient_adapter = GradientBasedAdapter::new(&config)?;
let performance_adapter = PerformanceBasedAdapter::new(&config)?;
let drift_adapter = DriftAwareAdapter::new(&config)?;
let resource_adapter = ResourceAwareAdapter::new(&config)?;
let meta_optimizer = MetaOptimizer::new(&config)?;
Ok(Self {
current_lr: config.base_lr,
base_lr: config.base_lr,
adaptation_strategy,
gradient_adapter,
performance_adapter,
drift_adapter,
resource_adapter,
meta_optimizer,
adaptation_history: VecDeque::with_capacity(config.history_window_size),
config,
})
}
pub fn record_memory_usage_mb(&mut self, usage_mb: f64) {
self.resource_adapter
.record_memory_usage(usage_mb, self.config.memory_budget_mb);
}
pub fn record_energy_sample(&mut self, joules: f64) {
self.resource_adapter.record_energy(joules);
}
pub fn record_throughput(&mut self, samples_per_second: A) {
self.resource_adapter.record_throughput(samples_per_second);
}
pub fn add_source_task(&mut self, task: TaskData<A>) {
self.meta_optimizer.add_source_task(task);
}
pub fn update_learning_rate(
&mut self,
gradients: &Array1<A>,
loss: A,
metrics: &HashMap<String, A>,
step: usize,
) -> Result<A> {
let step_started = Instant::now();
let frequency = self.config.adaptation_frequency.max(1);
if !step.is_multiple_of(frequency) {
self.gradient_adapter.observe_only(gradients);
self.performance_adapter.observe_only(loss);
self.resource_adapter
.record_step_time(step_started.elapsed(), self.config.step_time_budget);
return Ok(self.current_lr);
}
let mut signals = Vec::new();
if self.config.enable_gradient_adaptation {
if let Ok(signal) = self.gradient_adapter.generate_signal(gradients, step) {
signals.push(signal);
}
}
if self.config.enable_performance_adaptation {
if let Ok(signal) = self
.performance_adapter
.generate_signal(loss, metrics, step)
{
signals.push(signal);
}
}
if self.config.enable_drift_adaptation {
if let Ok(signal) = self.drift_adapter.generate_signal(gradients, step) {
signals.push(signal);
}
}
if self.config.enable_resource_adaptation {
if let Ok(signal) = self.resource_adapter.generate_signal(step) {
signals.push(signal);
}
}
let previous_lr = self.current_lr;
let decision = self.adaptation_strategy.resolve_signals(
signals,
previous_lr,
self.config.adaptation_sensitivity,
self.config.use_ensemble_voting,
step,
)?;
if self.config.enable_meta_learning {
let meta_adjustment = self.meta_optimizer.meta_optimize(&decision, step)?;
self.current_lr = self.apply_meta_adjustment(decision.new_lr, meta_adjustment);
} else {
self.current_lr = decision.new_lr;
}
self.current_lr = self
.current_lr
.clamp(self.config.min_lr, self.config.max_lr);
let event = AdaptationEvent {
timestamp: Instant::now(),
old_lr: previous_lr,
new_lr: self.current_lr,
trigger_signals: decision.contributing_signals,
effectiveness_score: None, };
self.adaptation_history.push_back(event);
if self.adaptation_history.len() > self.config.history_window_size {
self.adaptation_history.pop_front();
}
self.resource_adapter
.record_step_time(step_started.elapsed(), self.config.step_time_budget);
Ok(self.current_lr)
}
pub fn previous_lr(&self) -> Option<A> {
self.adaptation_history.back().map(|event| event.old_lr)
}
pub fn last_decision(&self) -> Option<&AdaptationDecision<A>> {
self.adaptation_strategy.last_decision.as_ref()
}
pub fn voting_history(&self) -> &VecDeque<SignalVote<A>> {
&self.adaptation_strategy.voting_history
}
pub fn get_current_lr(&self) -> A {
self.current_lr
}
pub fn get_adaptation_statistics(&self) -> AdaptationStatistics<A> {
let total_adaptations = self.adaptation_history.len();
let successful_adaptations = self
.adaptation_history
.iter()
.filter(|event| {
event
.effectiveness_score
.is_some_and(|score| score > A::zero())
})
.count();
let lr_volatility = if !self.adaptation_history.is_empty() {
let lr_values: Vec<A> = self
.adaptation_history
.iter()
.map(|event| event.new_lr)
.collect();
let mean_lr = lr_values.iter().fold(A::zero(), |acc, &lr| acc + lr)
/ scalar_or(lr_values.len(), A::one());
let variance = lr_values
.iter()
.map(|&lr| {
let diff = lr - mean_lr;
diff * diff
})
.fold(A::zero(), |acc, var| acc + var)
/ scalar_or(lr_values.len(), A::one());
variance.sqrt()
} else {
A::zero()
};
let signal_reliability_scores = self.adaptation_strategy.signal_reliability.clone();
let mut signal_effectiveness: HashMap<AdaptationSignalType, A> = HashMap::new();
let mut signal_counts: HashMap<AdaptationSignalType, usize> = HashMap::new();
for event in &self.adaptation_history {
let Some(score) = event.effectiveness_score else {
continue;
};
for signal_type in &event.trigger_signals {
let count = signal_counts.entry(*signal_type).or_insert(0);
*count += 1;
let steps = A::from(*count).unwrap_or_else(A::one);
let entry = signal_effectiveness
.entry(*signal_type)
.or_insert_with(A::zero);
*entry = *entry + (score - *entry) / steps;
}
}
let avg_adaptation_frequency = if let (Some(first), Some(last)) = (
self.adaptation_history.front(),
self.adaptation_history.back(),
) {
let span = last.timestamp.saturating_duration_since(first.timestamp);
if span.as_secs_f64() > 0.0 {
A::from(total_adaptations as f64 / span.as_secs_f64()).unwrap_or_else(A::zero)
} else {
A::zero()
}
} else {
A::zero()
};
let convergence_speed_improvement = {
let scores: Vec<A> = self
.adaptation_history
.iter()
.filter_map(|event| event.effectiveness_score)
.collect();
if scores.is_empty() {
A::zero()
} else {
let count = A::from(scores.len()).unwrap_or_else(A::one);
scores.iter().fold(A::zero(), |acc, score| acc + *score) / count
}
};
AdaptationStatistics {
total_adaptations,
successful_adaptations,
avg_adaptation_frequency,
lr_volatility,
signal_reliability_scores,
signal_effectiveness,
resource_efficiency_gains: A::from(self.resource_adapter.budget_violations() as f64)
.map(|violations| A::zero() - violations)
.unwrap_or_else(A::zero),
convergence_speed_improvement,
}
}
fn apply_meta_adjustment(&self, base_lr: A, meta_adjustment: A) -> A {
let alpha = scalar_or(0.7, A::zero()); let beta = scalar_or(0.3, A::zero());
alpha * base_lr + beta * meta_adjustment
}
pub fn evaluate_adaptation_effectiveness(&mut self, performance_improvement: A) {
let mut signals = Vec::new();
if let Some(last_event) = self.adaptation_history.back_mut() {
last_event.effectiveness_score = Some(performance_improvement);
signals = last_event.trigger_signals.clone();
}
for signal_type in signals {
self.adaptation_strategy
.update_signal_reliability(signal_type, performance_improvement);
}
self.meta_optimizer.record_reward(performance_improvement);
}
pub fn meta_arm_rewards(&self) -> &HashMap<usize, A> {
self.meta_optimizer.arm_rewards()
}
pub fn meta_arm_counts(&self) -> &HashMap<usize, usize> {
self.meta_optimizer.arm_counts()
}
pub fn signal_reliability(&self) -> &HashMap<AdaptationSignalType, A> {
&self.adaptation_strategy.signal_reliability
}
pub fn reset(&mut self) {
self.current_lr = self.base_lr;
self.adaptation_history.clear();
self.gradient_adapter.reset();
self.performance_adapter.reset();
self.drift_adapter.reset();
self.resource_adapter.reset();
self.meta_optimizer.reset();
}
}
mod meta;
mod signals;
#[cfg(test)]
mod adaptive_lr_tests;
pub(crate) use signals::LossDriftDetector;
impl<A: Float + Default + Send + Sync + Send + Sync> Default for GradientNormStatistics<A> {
fn default() -> Self {
Self {
mean: A::default(),
variance: A::default(),
skewness: A::default(),
kurtosis: A::default(),
percentiles: vec![A::default(); 5],
autocorrelation: A::default(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for SignalToNoiseEstimator<A> {
fn default() -> Self {
Self {
signal_estimate: A::default(),
noise_estimate: A::default(),
snr_history: VecDeque::new(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for PerformanceTrendAnalyzer<A> {
fn default() -> Self {
Self {
trend_detection_window: 10,
trend_types: vec![],
trend_strength: A::default(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for PlateauDetector<A> {
fn default() -> Self {
Self {
plateau_threshold: A::default(),
min_plateau_duration: 5,
current_plateau_length: 0,
plateau_confidence: A::default(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for OverfittingDetector<A> {
fn default() -> Self {
Self {
train_loss_history: VecDeque::new(),
val_loss_history: VecDeque::new(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for LearningEfficiencyTracker<A> {
fn default() -> Self {
Self {
loss_reduction_per_step: VecDeque::new(),
efficiency_score: A::default(),
efficiency_trend: TrendType::Improving,
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for DistributionTracker<A> {
fn default() -> Self {
Self {
feature_distributions: HashMap::new(),
distribution_drift_score: A::default(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for AdaptationSpeedController<A> {
fn default() -> Self {
Self {
base_adaptation_rate: A::from(0.1).unwrap_or_default(),
current_adaptation_rate: A::from(0.1).unwrap_or_default(),
acceleration_factor: A::from(1.1).unwrap_or_default(),
deceleration_factor: A::from(0.9).unwrap_or_default(),
momentum: A::default(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for DriftSeverityAssessor<A> {
fn default() -> Self {
Self {
severity_levels: vec![],
current_severity: DriftSeverityLevel::default(),
severity_history: VecDeque::new(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for DriftSeverityLevel<A> {
fn default() -> Self {
Self {
level: DriftSeverity::None,
recommended_lr_adjustment: A::one(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for ExplorationStrategy<A> {
fn default() -> Self {
Self {
exploration_rate: A::from(0.1).unwrap_or_default(),
arm_rewards: HashMap::new(),
arm_counts: HashMap::new(),
}
}
}
impl<A: Float + Default + Send + Sync + Send + Sync> Default for TransferLearner<A> {
fn default() -> Self {
Self {
source_task_data: vec![],
transfer_confidence: A::default(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::Array1;
#[test]
fn test_enhanced_adaptive_lr_controller_creation() {
let config = AdaptiveLRConfig {
base_lr: 0.01,
min_lr: 1e-6,
max_lr: 1.0,
enable_gradient_adaptation: true,
enable_performance_adaptation: true,
enable_drift_adaptation: false,
enable_resource_adaptation: false,
enable_meta_learning: false,
history_window_size: 100,
adaptation_frequency: 10,
adaptation_sensitivity: 0.1,
use_ensemble_voting: true,
step_time_budget: None,
memory_budget_mb: None,
};
let controller = EnhancedAdaptiveLRController::<f32>::new(config);
assert!(controller.is_ok());
}
#[test]
fn test_learning_rate_update() {
let config = AdaptiveLRConfig {
base_lr: 0.01,
min_lr: 1e-6,
max_lr: 1.0,
enable_gradient_adaptation: true,
enable_performance_adaptation: true,
enable_drift_adaptation: false,
enable_resource_adaptation: false,
enable_meta_learning: false,
history_window_size: 100,
adaptation_frequency: 10,
adaptation_sensitivity: 0.1,
use_ensemble_voting: true,
step_time_budget: None,
memory_budget_mb: None,
};
let mut controller =
EnhancedAdaptiveLRController::<f32>::new(config).expect("unwrap failed");
let gradients = Array1::from_vec(vec![0.1, 0.2, 0.05]);
let loss = 0.5;
let metrics = HashMap::new();
let new_lr = controller.update_learning_rate(&gradients, loss, &metrics, 1);
assert!(new_lr.is_ok());
assert!(new_lr.expect("unwrap failed") > 0.0);
}
#[test]
fn test_adaptation_statistics() {
let config = AdaptiveLRConfig {
base_lr: 0.01,
min_lr: 1e-6,
max_lr: 1.0,
enable_gradient_adaptation: true,
enable_performance_adaptation: true,
enable_drift_adaptation: false,
enable_resource_adaptation: false,
enable_meta_learning: false,
history_window_size: 100,
adaptation_frequency: 10,
adaptation_sensitivity: 0.1,
use_ensemble_voting: true,
step_time_budget: None,
memory_budget_mb: None,
};
let controller = EnhancedAdaptiveLRController::<f32>::new(config).expect("unwrap failed");
let stats = controller.get_adaptation_statistics();
assert_eq!(stats.total_adaptations, 0);
assert_eq!(stats.successful_adaptations, 0);
}
}