1use crate::{Optimizer, OptimizerError, OptimizerResult};
14use parking_lot::RwLock;
15use std::collections::HashMap;
16use std::sync::Arc;
17use torsh_core::{
18 device::{CpuDevice, DeviceType},
19 DType,
20};
21use torsh_tensor::{creation::randn, Tensor};
22
23#[derive(Debug, Clone)]
25pub struct NeuralOptimizerConfig {
26 pub meta_learning_rate: f32,
28 pub hidden_size: usize,
30 pub num_layers: usize,
32 pub device: Arc<CpuDevice>,
34 pub max_grad_norm: f32,
36 pub coordinate_wise: bool,
38 pub history_length: usize,
40}
41
42impl Default for NeuralOptimizerConfig {
43 fn default() -> Self {
44 Self {
45 meta_learning_rate: 0.001,
46 hidden_size: 20,
47 num_layers: 2,
48 device: Arc::new(CpuDevice::new()),
49 max_grad_norm: 10.0,
50 coordinate_wise: true,
51 history_length: 20,
52 }
53 }
54}
55
56#[derive(Debug, Clone)]
59pub struct OptimizerNetwork {
60 pub hidden_states: HashMap<String, Tensor>,
62 pub cell_states: HashMap<String, Tensor>,
64 pub weights: NetworkWeights,
66 pub config: NeuralOptimizerConfig,
68}
69
70#[derive(Debug, Clone)]
72pub struct NetworkWeights {
73 pub w_input: Tensor,
75 pub w_forget: Tensor,
77 pub w_output: Tensor,
79 pub w_cell: Tensor,
81 pub w_output_proj: Tensor,
83 pub bias_input: Tensor,
85 pub bias_forget: Tensor,
86 pub bias_output: Tensor,
87 pub bias_cell: Tensor,
88 pub bias_output_proj: Tensor,
89}
90
91impl NetworkWeights {
92 pub fn new(input_size: usize, hidden_size: usize, device: &CpuDevice) -> OptimizerResult<Self> {
94 let scale = (2.0 / (input_size + hidden_size) as f32).sqrt();
95
96 Ok(Self {
97 w_input: randn::<f32>(&[input_size + hidden_size, hidden_size])?.mul_scalar(scale)?,
98 w_forget: randn::<f32>(&[input_size + hidden_size, hidden_size])?.mul_scalar(scale)?,
99 w_output: randn::<f32>(&[input_size + hidden_size, hidden_size])?.mul_scalar(scale)?,
100 w_cell: randn::<f32>(&[input_size + hidden_size, hidden_size])?.mul_scalar(scale)?,
101 w_output_proj: randn::<f32>(&[hidden_size, 1])?.mul_scalar(scale)?,
102 bias_input: Tensor::zeros(&[hidden_size], DeviceType::Cpu)?,
103 bias_forget: Tensor::ones(&[hidden_size], DeviceType::Cpu)?, bias_output: Tensor::zeros(&[hidden_size], DeviceType::Cpu)?,
105 bias_cell: Tensor::zeros(&[hidden_size], DeviceType::Cpu)?,
106 bias_output_proj: Tensor::zeros(&[1], DeviceType::Cpu)?,
107 })
108 }
109}
110
111impl OptimizerNetwork {
112 pub fn new(config: NeuralOptimizerConfig) -> OptimizerResult<Self> {
114 let input_size = if config.coordinate_wise {
115 2 } else {
117 config.history_length * 2 };
119
120 let weights = NetworkWeights::new(input_size, config.hidden_size, &config.device)?;
121
122 Ok(Self {
123 hidden_states: HashMap::new(),
124 cell_states: HashMap::new(),
125 weights,
126 config,
127 })
128 }
129
130 pub fn forward(
132 &mut self,
133 param_id: &str,
134 gradient: &Tensor,
135 parameter: &Tensor,
136 ) -> OptimizerResult<Tensor> {
137 let device = self.config.device.clone();
138
139 let input = if self.config.coordinate_wise {
141 let grad_norm = gradient.norm()?.unsqueeze(0)?;
143 let param_norm = parameter.norm()?.unsqueeze(0)?;
144 Tensor::cat(&[&grad_norm, ¶m_norm], 0)?
145 } else {
146 let grad_flat = gradient.flatten()?;
148 let param_flat = parameter.flatten()?;
149 Tensor::cat(&[&grad_flat, ¶m_flat], 0)?
150 };
151
152 let hidden_shape = vec![self.config.hidden_size];
154 let hidden_state = self
155 .hidden_states
156 .entry(param_id.to_string())
157 .or_insert_with(|| {
158 Tensor::zeros(&hidden_shape, DeviceType::Cpu)
159 .expect("tensor creation should succeed")
160 });
161 let cell_state = self
162 .cell_states
163 .entry(param_id.to_string())
164 .or_insert_with(|| {
165 Tensor::zeros(&hidden_shape, DeviceType::Cpu)
166 .expect("tensor creation should succeed")
167 })
168 .clone();
169
170 let combined_input = Tensor::cat(&[&input, &hidden_state.clone()], 0)?;
172
173 let input_gate = self.sigmoid(
175 &combined_input
176 .matmul(&self.weights.w_input)?
177 .add_op(&self.weights.bias_input)?,
178 )?;
179 let forget_gate = self.sigmoid(
180 &combined_input
181 .matmul(&self.weights.w_forget)?
182 .add_op(&self.weights.bias_forget)?,
183 )?;
184 let output_gate = self.sigmoid(
185 &combined_input
186 .matmul(&self.weights.w_output)?
187 .add_op(&self.weights.bias_output)?,
188 )?;
189 let cell_gate = self.tanh(
190 &combined_input
191 .matmul(&self.weights.w_cell)?
192 .add_op(&self.weights.bias_cell)?,
193 )?;
194
195 let new_cell_state = forget_gate
197 .mul_op(&cell_state)?
198 .add_op(&input_gate.mul_op(&cell_gate)?)?;
199
200 let new_hidden_state = output_gate.mul_op(&self.tanh(&new_cell_state)?)?;
202
203 let update_magnitude = new_hidden_state
205 .matmul(&self.weights.w_output_proj)?
206 .add_op(&self.weights.bias_output_proj)?;
207
208 let update = if self.config.coordinate_wise {
210 gradient.mul_op(&update_magnitude.broadcast_to(gradient.shape().dims())?)?
212 } else {
213 gradient.mul_scalar(update_magnitude.item()?)?
215 };
216
217 *self
219 .hidden_states
220 .get_mut(param_id)
221 .expect("hidden_states should exist for param_id") = new_hidden_state;
222 *self
223 .cell_states
224 .get_mut(param_id)
225 .expect("cell_states should exist for param_id") = new_cell_state;
226
227 Ok(update)
228 }
229
230 fn sigmoid(&self, x: &Tensor) -> OptimizerResult<Tensor> {
232 let neg_x = x.mul_scalar(-1.0)?;
234 let exp_neg_x = neg_x.exp()?;
235 let one_plus_exp = exp_neg_x.add_scalar(1.0)?;
236 Ok(one_plus_exp.reciprocal()?)
237 }
238
239 fn tanh(&self, x: &Tensor) -> OptimizerResult<Tensor> {
241 Ok(x.tanh()?)
242 }
243
244 pub fn reset_state(&mut self) {
246 self.hidden_states.clear();
247 self.cell_states.clear();
248 }
249
250 pub fn parameters(&self) -> Vec<&Tensor> {
252 vec![
253 &self.weights.w_input,
254 &self.weights.w_forget,
255 &self.weights.w_output,
256 &self.weights.w_cell,
257 &self.weights.w_output_proj,
258 &self.weights.bias_input,
259 &self.weights.bias_forget,
260 &self.weights.bias_output,
261 &self.weights.bias_cell,
262 &self.weights.bias_output_proj,
263 ]
264 }
265}
266
267pub struct NeuralOptimizer {
269 pub network: OptimizerNetwork,
271 pub parameters: Vec<Tensor>,
273 pub meta_optimizer: Option<Box<dyn Optimizer>>,
275 pub training: bool,
277 pub step_count: usize,
279}
280
281impl NeuralOptimizer {
282 pub fn new(
284 parameters: Vec<Tensor>,
285 config: Option<NeuralOptimizerConfig>,
286 ) -> OptimizerResult<Self> {
287 let config = config.unwrap_or_default();
288 let network = OptimizerNetwork::new(config)?;
289
290 Ok(Self {
291 network,
292 parameters,
293 meta_optimizer: None,
294 training: false,
295 step_count: 0,
296 })
297 }
298
299 pub fn with_meta_learning(
301 parameters: Vec<Tensor>,
302 config: Option<NeuralOptimizerConfig>,
303 ) -> OptimizerResult<Self> {
304 let mut optimizer = Self::new(parameters, config)?;
305 optimizer.training = true;
306
307 let network_params = optimizer
309 .network
310 .parameters()
311 .iter()
312 .map(|p| Arc::new(RwLock::new((*p).clone())))
313 .collect();
314
315 use crate::adam::Adam;
316 let meta_optimizer = Adam::new(
317 network_params,
318 Some(optimizer.network.config.meta_learning_rate),
319 None,
320 None,
321 None,
322 false,
323 );
324
325 optimizer.meta_optimizer = Some(Box::new(meta_optimizer));
326
327 Ok(optimizer)
328 }
329
330 pub fn train(&mut self, mode: bool) {
332 self.training = mode;
333 }
334
335 pub fn reset(&mut self) {
337 self.network.reset_state();
338 self.step_count = 0;
339 }
340
341 pub fn compute_meta_loss(&self, target_loss: f32, actual_loss: f32) -> f32 {
343 (target_loss - actual_loss).powi(2)
344 }
345
346 pub fn meta_step(&mut self, meta_loss: f32) -> OptimizerResult<()> {
348 if let Some(ref mut meta_optimizer) = self.meta_optimizer {
349 for param in self.network.parameters() {
352 let meta_grad =
354 randn::<f32>(param.shape().dims())?.mul_scalar(meta_loss * 0.001)?;
355 param.set_grad(Some(meta_grad));
356 }
357
358 meta_optimizer.step()?;
359 }
360 Ok(())
361 }
362}
363
364impl Optimizer for NeuralOptimizer {
365 fn step(&mut self) -> OptimizerResult<()> {
366 self.step_count += 1;
367
368 for (i, param) in self.parameters.iter_mut().enumerate() {
369 if let Some(grad) = param.grad() {
370 let param_id = format!("param_{}", i);
371
372 let update = self.network.forward(¶m_id, &grad, param)?;
374
375 let update_norm = update.norm()?.item()?;
377 let clipped_update = if update_norm > self.network.config.max_grad_norm {
378 update.mul_scalar(self.network.config.max_grad_norm / update_norm)?
379 } else {
380 update
381 };
382
383 crate::param_update::sub_assign(&mut *param, &clipped_update)?;
385
386 param.set_grad(None);
388 }
389 }
390
391 Ok(())
392 }
393
394 fn zero_grad(&mut self) {
395 for param in &mut self.parameters {
396 }
399 }
400
401 fn get_lr(&self) -> Vec<f32> {
402 vec![self.network.config.meta_learning_rate]
404 }
405
406 fn set_lr(&mut self, lr: f32) {
407 self.network.config.meta_learning_rate = lr;
409 }
410
411 fn state_dict(&self) -> OptimizerResult<crate::OptimizerState> {
412 let mut state = crate::OptimizerState::new("NeuralOptimizer".to_string());
413
414 state.global_state.insert(
416 "meta_learning_rate".to_string(),
417 self.network.config.meta_learning_rate,
418 );
419 state
420 .global_state
421 .insert("step_count".to_string(), self.step_count as f32);
422 state.global_state.insert(
423 "hidden_size".to_string(),
424 self.network.config.hidden_size as f32,
425 );
426 state.global_state.insert(
427 "num_layers".to_string(),
428 self.network.config.num_layers as f32,
429 );
430
431 Ok(state)
434 }
435
436 fn add_param_group(
437 &mut self,
438 params: Vec<std::sync::Arc<parking_lot::RwLock<Tensor>>>,
439 options: std::collections::HashMap<String, f32>,
440 ) {
441 }
444
445 fn load_state_dict(&mut self, state: crate::OptimizerState) -> OptimizerResult<()> {
446 if let Some(&meta_lr) = state.global_state.get("meta_learning_rate") {
447 self.network.config.meta_learning_rate = meta_lr;
448 }
449
450 if let Some(&step_count) = state.global_state.get("step_count") {
451 self.step_count = step_count as usize;
452 }
453
454 Ok(())
455 }
456}
457
458pub struct NeuralOptimizerTrainer {
460 pub optimizer: NeuralOptimizer,
462 pub training_tasks: Vec<Box<dyn OptimizationTask>>,
464 pub validation_tasks: Vec<Box<dyn OptimizationTask>>,
466 pub config: TrainingConfig,
468}
469
470#[derive(Debug, Clone)]
472pub struct TrainingConfig {
473 pub meta_iterations: usize,
475 pub inner_steps: usize,
477 pub meta_lr: f32,
479 pub device: Arc<CpuDevice>,
481}
482
483impl Default for TrainingConfig {
484 fn default() -> Self {
485 Self {
486 meta_iterations: 1000,
487 inner_steps: 100,
488 meta_lr: 0.001,
489 device: Arc::new(CpuDevice::new()),
490 }
491 }
492}
493
494pub trait OptimizationTask {
496 fn initialize_parameters(&self, device: &CpuDevice) -> OptimizerResult<Vec<Tensor>>;
498
499 fn compute_loss_and_gradients(
501 &self,
502 parameters: &[Tensor],
503 ) -> OptimizerResult<(f32, Vec<Tensor>)>;
504
505 fn name(&self) -> &str;
507}
508
509pub struct QuadraticTask {
511 pub dimension: usize,
512 pub condition_number: f32,
513 pub name: String,
514}
515
516impl QuadraticTask {
517 pub fn new(dimension: usize, condition_number: f32) -> Self {
518 Self {
519 dimension,
520 condition_number,
521 name: format!("Quadratic_{}D_cond{:.1}", dimension, condition_number),
522 }
523 }
524}
525
526impl OptimizationTask for QuadraticTask {
527 fn initialize_parameters(&self, device: &CpuDevice) -> OptimizerResult<Vec<Tensor>> {
528 Ok(vec![randn::<f32>(&[self.dimension])?])
529 }
530
531 fn compute_loss_and_gradients(
532 &self,
533 parameters: &[Tensor],
534 ) -> OptimizerResult<(f32, Vec<Tensor>)> {
535 let param = ¶meters[0];
536
537 let mut hessian_diag = Vec::new();
540 for i in 0..self.dimension {
541 let eigenval =
542 1.0 + (self.condition_number - 1.0) * (i as f32) / (self.dimension as f32 - 1.0);
543 hessian_diag.push(eigenval);
544 }
545
546 let hessian_diag_tensor =
547 Tensor::from_data(hessian_diag, param.shape().dims().to_vec(), param.device())?;
548
549 let loss = param
551 .pow(2.0)?
552 .mul_op(&hessian_diag_tensor)?
553 .sum()?
554 .mul_scalar(0.5)?
555 .item()?;
556
557 let grad = param.mul_op(&hessian_diag_tensor)?;
559
560 Ok((loss, vec![grad]))
561 }
562
563 fn name(&self) -> &str {
564 &self.name
565 }
566}
567
568impl NeuralOptimizerTrainer {
569 pub fn new(
571 optimizer: NeuralOptimizer,
572 training_tasks: Vec<Box<dyn OptimizationTask>>,
573 config: Option<TrainingConfig>,
574 ) -> Self {
575 Self {
576 optimizer,
577 training_tasks,
578 validation_tasks: Vec::new(),
579 config: config.unwrap_or_default(),
580 }
581 }
582
583 pub fn train(&mut self) -> OptimizerResult<Vec<f32>> {
585 let mut meta_losses = Vec::new();
586
587 for meta_iter in 0..self.config.meta_iterations {
588 let mut total_meta_loss = 0.0;
589
590 let task_idx = meta_iter % self.training_tasks.len();
592 let task = &self.training_tasks[task_idx];
593
594 let mut params = task.initialize_parameters(&CpuDevice::default())?;
596
597 let mut task_loss = 0.0;
599 for _ in 0..self.config.inner_steps {
600 let (loss, grads) = task.compute_loss_and_gradients(¶ms)?;
601 task_loss = loss;
602
603 for (param, grad) in params.iter_mut().zip(grads.iter()) {
605 param.set_grad(Some(grad.clone()));
606 }
607
608 self.optimizer.step()?;
611 }
612
613 let target_loss = 0.0; let meta_loss = self.optimizer.compute_meta_loss(target_loss, task_loss);
616 total_meta_loss += meta_loss;
617
618 self.optimizer.meta_step(meta_loss)?;
620
621 meta_losses.push(total_meta_loss);
622
623 if meta_iter % 100 == 0 {
624 println!(
625 "Meta-iteration {}: Meta-loss = {:.6}, Task loss = {:.6} (Task: {})",
626 meta_iter,
627 meta_loss,
628 task_loss,
629 task.name()
630 );
631 }
632 }
633
634 Ok(meta_losses)
635 }
636}
637
638#[cfg(test)]
639mod tests {
640 use super::*;
641
642 #[test]
643 fn test_neural_optimizer_config() {
644 let config = NeuralOptimizerConfig::default();
645 assert_eq!(config.meta_learning_rate, 0.001);
646 assert_eq!(config.hidden_size, 20);
647 assert_eq!(config.num_layers, 2);
648 assert_eq!(config.max_grad_norm, 10.0);
649 assert!(config.coordinate_wise);
650 assert_eq!(config.history_length, 20);
651 }
652
653 #[test]
654 fn test_quadratic_task() {
655 let task = QuadraticTask::new(10, 100.0);
656 assert_eq!(task.dimension, 10);
657 assert_eq!(task.condition_number, 100.0);
658 assert_eq!(task.name(), "Quadratic_10D_cond100.0");
659 }
660
661 #[test]
662 fn test_training_config() {
663 let config = TrainingConfig::default();
664 assert_eq!(config.meta_iterations, 1000);
665 assert_eq!(config.inner_steps, 100);
666 assert_eq!(config.meta_lr, 0.001);
667 }
668}