use crate::{Adam, Optimizer, OptimizerError, OptimizerResult, SGD};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use torsh_tensor::Tensor;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TaskCharacteristics {
pub dimension: usize,
pub training_steps: usize,
pub problem_type: ProblemType,
pub gradient_stats: GradientStatistics,
pub landscape_properties: LandscapeProperties,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ProblemType {
Classification,
Regression,
Reinforcement,
Generative,
Custom(String),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct GradientStatistics {
pub mean_magnitude: f32,
pub variance_magnitude: f32,
pub sparsity_ratio: f32,
pub correlation_length: f32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LandscapeProperties {
pub estimated_smoothness: f32,
pub condition_number_estimate: f32,
pub has_saddle_points: bool,
pub convexity_score: f32,
}
#[derive(Debug, Clone)]
pub enum MetaLearningAlgorithm {
MAML {
inner_lr: f32,
outer_lr: f32,
inner_steps: usize,
},
L2O {
meta_optimizer: Box<dyn Optimizer>,
hidden_size: usize,
},
FewShot {
adaptation_steps: usize,
adaptation_lr: f32,
},
GradientBased {
meta_step_size: f32,
adaptation_steps: usize,
},
}
#[derive(Debug, Clone)]
pub struct MetaLearningConfig {
pub algorithm: MetaLearningAlgorithm,
pub num_meta_tasks: usize,
pub support_set_size: usize,
pub query_set_size: usize,
pub meta_batch_size: usize,
pub meta_epochs: usize,
}
pub struct TaskDataset {
pub tasks: Vec<Task>,
pub meta_split: (Vec<usize>, Vec<usize>), }
#[derive(Debug, Clone)]
pub struct Task {
pub id: String,
pub characteristics: TaskCharacteristics,
pub support_data: Vec<(Tensor, Tensor)>, pub query_data: Vec<(Tensor, Tensor)>,
pub metadata: HashMap<String, String>,
}
pub struct MetaOptimizer {
config: MetaLearningConfig,
base_optimizer: Box<dyn Optimizer>,
meta_parameters: HashMap<String, Tensor>,
task_history: Vec<TaskPerformance>,
adaptation_rules: AdaptationRules,
}
#[derive(Debug, Clone)]
pub struct TaskPerformance {
pub task_id: String,
pub characteristics: TaskCharacteristics,
pub initial_loss: f32,
pub final_loss: f32,
pub convergence_steps: usize,
pub optimizer_config: OptimizerConfig,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OptimizerConfig {
pub optimizer_type: String,
pub hyperparameters: HashMap<String, f32>,
}
#[derive(Debug, Clone)]
pub struct AdaptationRules {
pub characteristic_rules: Vec<(TaskMatcher, OptimizerConfig)>,
pub learned_patterns: HashMap<String, Vec<f32>>,
pub performance_thresholds: HashMap<String, f32>,
}
#[derive(Debug, Clone)]
pub struct TaskMatcher {
pub dimension_range: Option<(usize, usize)>,
pub problem_type: Option<ProblemType>,
pub gradient_magnitude_range: Option<(f32, f32)>,
pub sparsity_range: Option<(f32, f32)>,
}
impl MetaOptimizer {
pub fn new(config: MetaLearningConfig, base_optimizer: Box<dyn Optimizer>) -> Self {
Self {
config,
base_optimizer,
meta_parameters: HashMap::new(),
task_history: Vec::new(),
adaptation_rules: AdaptationRules::new(),
}
}
pub fn meta_train(&mut self, task_dataset: &TaskDataset) -> OptimizerResult<()> {
match &self.config.algorithm {
MetaLearningAlgorithm::MAML { inner_lr, outer_lr, inner_steps } => {
self.maml_training(task_dataset, *inner_lr, *outer_lr, *inner_steps)
}
MetaLearningAlgorithm::L2O { .. } => {
self.l2o_training(task_dataset)
}
MetaLearningAlgorithm::FewShot { adaptation_steps, adaptation_lr } => {
self.few_shot_training(task_dataset, *adaptation_steps, *adaptation_lr)
}
MetaLearningAlgorithm::GradientBased { meta_step_size, adaptation_steps } => {
self.gradient_based_training(task_dataset, *meta_step_size, *adaptation_steps)
}
}
}
pub fn adapt_to_task(&mut self, task: &Task) -> OptimizerResult<Box<dyn Optimizer>> {
let characteristics = &task.characteristics;
let optimizer_config = self.select_optimizer_config(characteristics)?;
let adapted_optimizer = self.create_optimizer_from_config(&optimizer_config)?;
let final_optimizer = match &self.config.algorithm {
MetaLearningAlgorithm::FewShot { adaptation_steps, adaptation_lr } => {
self.perform_few_shot_adaptation(adapted_optimizer, task, *adaptation_steps, *adaptation_lr)?
}
_ => adapted_optimizer,
};
Ok(final_optimizer)
}
fn maml_training(&mut self, task_dataset: &TaskDataset, inner_lr: f32, outer_lr: f32, inner_steps: usize) -> OptimizerResult<()> {
for meta_epoch in 0..self.config.meta_epochs {
let mut meta_gradients = HashMap::new();
let task_batch = self.sample_task_batch(task_dataset, self.config.meta_batch_size)?;
for task_id in task_batch {
let task = &task_dataset.tasks[task_id];
let mut adapted_params = self.meta_parameters.clone();
for _ in 0..inner_steps {
let support_gradients = self.compute_task_gradients(task, &adapted_params, true)?;
for (param_name, gradient) in support_gradients {
if let Some(param) = adapted_params.get_mut(¶m_name) {
*param = param.sub(&gradient.mul_scalar(inner_lr)?)?;
}
}
}
let query_gradients = self.compute_task_gradients(task, &adapted_params, false)?;
for (param_name, gradient) in query_gradients {
match meta_gradients.get_mut(¶m_name) {
Some(existing) => {
*existing = existing.add(&gradient)?;
}
None => {
meta_gradients.insert(param_name, gradient);
}
}
}
}
for (param_name, meta_gradient) in meta_gradients {
if let Some(param) = self.meta_parameters.get_mut(¶m_name) {
*param = param.sub(&meta_gradient.mul_scalar(outer_lr)?)?;
}
}
log::info!("Meta-epoch {}/{} completed", meta_epoch + 1, self.config.meta_epochs);
}
Ok(())
}
fn l2o_training(&mut self, task_dataset: &TaskDataset) -> OptimizerResult<()> {
for meta_epoch in 0..self.config.meta_epochs {
let task_batch = self.sample_task_batch(task_dataset, self.config.meta_batch_size)?;
for task_id in task_batch {
let task = &task_dataset.tasks[task_id];
let state_features = self.extract_optimizer_state_features(task)?;
let predicted_update = self.predict_optimizer_update(&state_features)?;
let performance = self.evaluate_optimizer_update(task, &predicted_update)?;
self.update_meta_optimizer(performance)?;
}
}
Ok(())
}
fn few_shot_training(&mut self, task_dataset: &TaskDataset, adaptation_steps: usize, _adaptation_lr: f32) -> OptimizerResult<()> {
for task in &task_dataset.tasks {
let configs = self.generate_optimizer_configs(&task.characteristics)?;
for config in configs {
let performance = self.evaluate_config_on_task(&config, task, adaptation_steps)?;
self.task_history.push(TaskPerformance {
task_id: task.id.clone(),
characteristics: task.characteristics.clone(),
initial_loss: performance.initial_loss,
final_loss: performance.final_loss,
convergence_steps: performance.convergence_steps,
optimizer_config: config,
});
}
}
self.learn_adaptation_rules()?;
Ok(())
}
fn gradient_based_training(&mut self, task_dataset: &TaskDataset, meta_step_size: f32, adaptation_steps: usize) -> OptimizerResult<()> {
for meta_epoch in 0..self.config.meta_epochs {
let task_batch = self.sample_task_batch(task_dataset, self.config.meta_batch_size)?;
for task_id in task_batch {
let task = &task_dataset.tasks[task_id];
let meta_gradients = self.compute_meta_gradients(task, adaptation_steps)?;
for (param_name, gradient) in meta_gradients {
if let Some(param) = self.meta_parameters.get_mut(¶m_name) {
*param = param.sub(&gradient.mul_scalar(meta_step_size)?)?;
}
}
}
}
Ok(())
}
fn select_optimizer_config(&self, characteristics: &TaskCharacteristics) -> OptimizerResult<OptimizerConfig> {
for (matcher, config) in &self.adaptation_rules.characteristic_rules {
if matcher.matches(characteristics) {
return Ok(config.clone());
}
}
Ok(OptimizerConfig {
optimizer_type: "Adam".to_string(),
hyperparameters: {
let mut params = HashMap::new();
params.insert("lr".to_string(), 0.001);
params.insert("beta1".to_string(), 0.9);
params.insert("beta2".to_string(), 0.999);
params
},
})
}
fn create_optimizer_from_config(&self, config: &OptimizerConfig) -> OptimizerResult<Box<dyn Optimizer>> {
let empty_params: Vec<Arc<RwLock<Tensor>>> = Vec::new();
match config.optimizer_type.to_lowercase().as_str() {
"adam" | "adamw" => {
let lr = config.hyperparameters.get("lr").copied();
let beta1 = config.hyperparameters.get("beta1").copied().unwrap_or(0.9);
let beta2 = config.hyperparameters.get("beta2").copied().unwrap_or(0.999);
let eps = config.hyperparameters.get("eps").copied();
let weight_decay = config.hyperparameters.get("weight_decay").copied();
Ok(Box::new(Adam::new(
empty_params,
lr,
Some((beta1, beta2)),
eps,
weight_decay,
false,
)))
}
"sgd" => {
let lr = config.hyperparameters.get("lr").copied().unwrap_or(0.01);
let momentum = config.hyperparameters.get("momentum").copied();
let weight_decay = config.hyperparameters.get("weight_decay").copied();
let dampening = config.hyperparameters.get("dampening").copied();
Ok(Box::new(SGD::new(
empty_params,
lr,
momentum,
dampening,
weight_decay,
false,
)))
}
other => Err(OptimizerError::ConfigError(format!(
"Unknown optimizer type: '{}'. Supported types: adam, adamw, sgd",
other
))),
}
}
fn perform_few_shot_adaptation(&self, mut optimizer: Box<dyn Optimizer>, _task: &Task, adaptation_steps: usize, _adaptation_lr: f32) -> OptimizerResult<Box<dyn Optimizer>> {
for _ in 0..adaptation_steps {
optimizer.step()?;
}
Ok(optimizer)
}
fn sample_task_batch(&self, task_dataset: &TaskDataset, batch_size: usize) -> Vec<usize> {
let mut batch = Vec::new();
let train_tasks = &task_dataset.meta_split.0;
for _ in 0..batch_size {
if !train_tasks.is_empty() {
let idx = fastrand::usize(0..train_tasks.len());
batch.push(train_tasks[idx]);
}
}
batch
}
fn compute_task_gradients(&self, task: &Task, _parameters: &HashMap<String, Tensor>, use_support: bool) -> OptimizerResult<HashMap<String, Tensor>> {
let split = if use_support { "support" } else { "query" };
Err(OptimizerError::ConfigError(format!(
"compute_task_gradients (task '{}', {} set) is not available: MetaOptimizer \
has no differentiable model/loss to backpropagate through. Provide a model \
and loss-aware training loop before invoking gradient-based meta-learning.",
task.id, split
)))
}
fn extract_optimizer_state_features(&self, task: &Task) -> OptimizerResult<Vec<f32>> {
let mut features = Vec::new();
features.push(task.characteristics.dimension as f32);
features.push(task.characteristics.gradient_stats.mean_magnitude);
features.push(task.characteristics.gradient_stats.variance_magnitude);
features.push(task.characteristics.gradient_stats.sparsity_ratio);
Ok(features)
}
fn predict_optimizer_update(&self, _state_features: &[f32]) -> OptimizerResult<HashMap<String, Tensor>> {
Err(OptimizerError::ConfigError(
"predict_optimizer_update is not available: the L2O meta-network that maps \
optimizer-state features to parameter updates has not been instantiated. \
A learned update model must be provided before L2O inference can run."
.to_string(),
))
}
fn evaluate_optimizer_update(&self, task: &Task, _update: &HashMap<String, Tensor>) -> OptimizerResult<f32> {
Err(OptimizerError::ConfigError(format!(
"evaluate_optimizer_update (task '{}') is not available: scoring an update \
requires a model and loss to measure post-update task performance.",
task.id
)))
}
fn update_meta_optimizer(&mut self, performance: f32) -> OptimizerResult<()> {
Ok(())
}
fn generate_optimizer_configs(&self, characteristics: &TaskCharacteristics) -> Vec<OptimizerConfig> {
let mut configs = Vec::new();
match characteristics.problem_type {
ProblemType::Classification => {
configs.push(OptimizerConfig {
optimizer_type: "Adam".to_string(),
hyperparameters: {
let mut params = HashMap::new();
params.insert("lr".to_string(), 0.001);
params.insert("beta1".to_string(), 0.9);
params.insert("beta2".to_string(), 0.999);
params
},
});
}
ProblemType::Regression => {
configs.push(OptimizerConfig {
optimizer_type: "SGD".to_string(),
hyperparameters: {
let mut params = HashMap::new();
params.insert("lr".to_string(), 0.01);
params.insert("momentum".to_string(), 0.9);
params
},
});
}
_ => {
configs.push(OptimizerConfig {
optimizer_type: "Adam".to_string(),
hyperparameters: HashMap::new(),
});
}
}
configs
}
fn evaluate_config_on_task(&self, _config: &OptimizerConfig, task: &Task, _max_steps: usize) -> OptimizerResult<TaskPerformance> {
Err(OptimizerError::ConfigError(format!(
"evaluate_config_on_task (task '{}') is not available: measuring real \
initial/final losses requires training a model with a loss function, \
which MetaOptimizer does not provide. Refusing to return fabricated metrics.",
task.id
)))
}
fn learn_adaptation_rules(&mut self) -> OptimizerResult<()> {
let mut rules = Vec::new();
let mut characteristic_groups: HashMap<String, Vec<&TaskPerformance>> = HashMap::new();
for performance in &self.task_history {
let key = format!("{:?}_{}",
performance.characteristics.problem_type,
performance.characteristics.dimension / 1000 );
characteristic_groups.entry(key).or_default().push(performance);
}
for (_, group) in characteristic_groups {
if let Some(best_performance) = group.iter().min_by(|a, b| a.final_loss.partial_cmp(&b.final_loss).unwrap_or(std::cmp::Ordering::Equal)) {
let matcher = TaskMatcher {
dimension_range: Some((best_performance.characteristics.dimension.saturating_sub(100), best_performance.characteristics.dimension + 100)),
problem_type: Some(best_performance.characteristics.problem_type.clone()),
gradient_magnitude_range: None,
sparsity_range: None,
};
rules.push((matcher, best_performance.optimizer_config.clone()));
}
}
self.adaptation_rules.characteristic_rules = rules;
Ok(())
}
fn compute_meta_gradients(&self, task: &Task, _adaptation_steps: usize) -> OptimizerResult<HashMap<String, Tensor>> {
Err(OptimizerError::ConfigError(format!(
"compute_meta_gradients (task '{}') is not available: higher-order \
meta-gradients require a differentiable model/loss and second-order \
autograd through the inner adaptation loop, which MetaOptimizer does not \
provide.",
task.id
)))
}
}
impl AdaptationRules {
pub fn new() -> Self {
Self {
characteristic_rules: Vec::new(),
learned_patterns: HashMap::new(),
performance_thresholds: HashMap::new(),
}
}
}
impl TaskMatcher {
pub fn matches(&self, characteristics: &TaskCharacteristics) -> bool {
if let Some((min_dim, max_dim)) = self.dimension_range {
if characteristics.dimension < min_dim || characteristics.dimension > max_dim {
return false;
}
}
if let Some(ref expected_type) = self.problem_type {
if std::mem::discriminant(&characteristics.problem_type) != std::mem::discriminant(expected_type) {
return false;
}
}
if let Some((min_mag, max_mag)) = self.gradient_magnitude_range {
if characteristics.gradient_stats.mean_magnitude < min_mag ||
characteristics.gradient_stats.mean_magnitude > max_mag {
return false;
}
}
if let Some((min_sparse, max_sparse)) = self.sparsity_range {
if characteristics.gradient_stats.sparsity_ratio < min_sparse ||
characteristics.gradient_stats.sparsity_ratio > max_sparse {
return false;
}
}
true
}
}
pub mod utils {
use super::*;
pub fn analyze_task_characteristics(gradients: &[Tensor]) -> OptimizerResult<TaskCharacteristics> {
if gradients.is_empty() {
return Err(OptimizerError::InvalidParameter("Empty gradient history".to_string()));
}
let mut magnitudes = Vec::new();
let mut sparsity_ratios = Vec::new();
for gradient in gradients {
let magnitude = gradient.norm()?.item()?;
magnitudes.push(magnitude);
let total_elements = gradient.numel()?;
let zero_threshold = 1e-8;
let non_zero_count = gradient.abs()?.gt_scalar(zero_threshold)?.sum()?.item()? as usize;
let sparsity = 1.0 - (non_zero_count as f32 / total_elements as f32);
sparsity_ratios.push(sparsity);
}
let mean_magnitude = magnitudes.iter().sum::<f32>() / magnitudes.len() as f32;
let variance_magnitude = magnitudes.iter()
.map(|x| (x - mean_magnitude).powi(2))
.sum::<f32>() / magnitudes.len() as f32;
let mean_sparsity = sparsity_ratios.iter().sum::<f32>() / sparsity_ratios.len() as f32;
Ok(TaskCharacteristics {
dimension: gradients[0].numel(),
training_steps: gradients.len(),
problem_type: ProblemType::Custom("unknown".to_string()),
gradient_stats: GradientStatistics {
mean_magnitude,
variance_magnitude,
sparsity_ratio: mean_sparsity,
correlation_length: 1.0, },
landscape_properties: LandscapeProperties {
estimated_smoothness: 0.5,
condition_number_estimate: 1.0,
has_saddle_points: false,
convexity_score: 0.5,
},
})
}
pub fn create_meta_dataset(tasks: Vec<Task>, train_ratio: f32) -> TaskDataset {
let n_train = (tasks.len() as f32 * train_ratio) as usize;
let train_ids: Vec<usize> = (0..n_train).collect();
let test_ids: Vec<usize> = (n_train..tasks.len()).collect();
TaskDataset {
tasks,
meta_split: (train_ids, test_ids),
}
}
pub fn evaluate_meta_performance(_meta_optimizer: &mut MetaOptimizer, _test_tasks: &[Task]) -> OptimizerResult<f32> {
Err(OptimizerError::ConfigError(
"evaluate_meta_performance is not available: scoring adapted optimizers on \
their query sets requires a model and loss function to measure real task \
performance. Refusing to return a fabricated score."
.to_string(),
))
}
}