Skip to main content

quantrs2_anneal/rl_embedding_optimizer/
networks.rs

1//! Neural network implementations for RL embedding optimization
2
3use scirs2_core::random::prelude::*;
4use scirs2_core::random::ChaCha8Rng;
5use scirs2_core::random::{Rng, SeedableRng};
6use std::collections::HashMap;
7use std::time::{Duration, Instant};
8
9use super::error::{RLEmbeddingError, RLEmbeddingResult};
10
11/// Deep Q-Network for embedding decisions
12#[derive(Debug, Clone)]
13pub struct EmbeddingDQN {
14    /// Main Q-network
15    pub q_network: EmbeddingNetwork,
16    /// Target Q-network for stable training
17    pub target_network: EmbeddingNetwork,
18    /// Network configuration
19    pub config: NetworkConfig,
20    /// Training state
21    pub training_state: NetworkTrainingState,
22}
23
24/// Policy network for continuous embedding optimization
25#[derive(Debug, Clone)]
26pub struct EmbeddingPolicyNetwork {
27    /// Actor network (policy)
28    pub actor_network: EmbeddingNetwork,
29    /// Critic network (value function)
30    pub critic_network: EmbeddingNetwork,
31    /// Network configuration
32    pub config: NetworkConfig,
33    /// Training state
34    pub training_state: NetworkTrainingState,
35}
36
37/// Neural network for embedding optimization
38#[derive(Debug, Clone)]
39pub struct EmbeddingNetwork {
40    /// Network layers
41    pub layers: Vec<NetworkLayer>,
42    /// Input normalization
43    pub input_norm: NormalizationLayer,
44    /// Output scaling
45    pub output_scaling: NormalizationLayer,
46    /// Network metadata
47    pub metadata: NetworkMetadata,
48}
49
50/// Neural network layer
51#[derive(Debug, Clone)]
52pub struct NetworkLayer {
53    /// Layer weights
54    pub weights: Vec<Vec<f64>>,
55    /// Layer biases
56    pub biases: Vec<f64>,
57    /// Activation function
58    pub activation: ActivationFunction,
59    /// Dropout rate
60    pub dropout_rate: f64,
61    /// Batch normalization parameters
62    pub batch_norm: Option<BatchNormalization>,
63}
64
65/// Activation functions
66#[derive(Debug, Clone, PartialEq)]
67pub enum ActivationFunction {
68    /// `ReLU` activation
69    ReLU,
70    /// Leaky `ReLU`
71    LeakyReLU(f64),
72    /// Tanh activation
73    Tanh,
74    /// Sigmoid activation
75    Sigmoid,
76    /// Swish activation
77    Swish,
78    /// Linear activation
79    Linear,
80}
81
82/// Batch normalization parameters
83#[derive(Debug, Clone)]
84pub struct BatchNormalization {
85    /// Running mean
86    pub running_mean: Vec<f64>,
87    /// Running variance
88    pub running_var: Vec<f64>,
89    /// Learnable scale parameter
90    pub gamma: Vec<f64>,
91    /// Learnable shift parameter
92    pub beta: Vec<f64>,
93    /// Epsilon for numerical stability
94    pub epsilon: f64,
95    /// Momentum for running statistics
96    pub momentum: f64,
97}
98
99/// Normalization layer
100#[derive(Debug, Clone)]
101pub struct NormalizationLayer {
102    /// Mean values
103    pub means: Vec<f64>,
104    /// Standard deviations
105    pub stds: Vec<f64>,
106    /// Min values
107    pub mins: Vec<f64>,
108    /// Max values
109    pub maxs: Vec<f64>,
110    /// Normalization type
111    pub norm_type: NormalizationType,
112}
113
114/// Types of normalization
115#[derive(Debug, Clone, PartialEq, Eq)]
116pub enum NormalizationType {
117    /// Z-score normalization
118    StandardScore,
119    /// Min-max normalization
120    MinMax,
121    /// Robust normalization (median, IQR)
122    Robust,
123    /// No normalization
124    None,
125}
126
127/// Network configuration
128#[derive(Debug, Clone)]
129pub struct NetworkConfig {
130    /// Layer sizes
131    pub layer_sizes: Vec<usize>,
132    /// Learning rate
133    pub learning_rate: f64,
134    /// Regularization parameters
135    pub regularization: RegularizationConfig,
136    /// Optimization method
137    pub optimizer: OptimizerType,
138    /// Loss function
139    pub loss_function: LossFunction,
140}
141
142/// Regularization configuration
143#[derive(Debug, Clone)]
144pub struct RegularizationConfig {
145    /// L1 regularization strength
146    pub l1_strength: f64,
147    /// L2 regularization strength
148    pub l2_strength: f64,
149    /// Dropout rate
150    pub dropout_rate: f64,
151    /// Early stopping patience
152    pub early_stopping_patience: usize,
153}
154
155/// Optimizer types
156#[derive(Debug, Clone, PartialEq)]
157pub enum OptimizerType {
158    /// Stochastic Gradient Descent
159    SGD,
160    /// Adam optimizer
161    Adam { beta1: f64, beta2: f64 },
162    /// `RMSprop` optimizer
163    RMSprop { decay_rate: f64 },
164    /// `AdaGrad` optimizer
165    AdaGrad,
166}
167
168/// Loss functions
169#[derive(Debug, Clone, PartialEq)]
170pub enum LossFunction {
171    /// Mean Squared Error
172    MSE,
173    /// Huber loss
174    Huber { delta: f64 },
175    /// Cross-entropy loss
176    CrossEntropy,
177    /// Custom multi-objective loss
178    MultiObjective,
179}
180
181/// Network metadata
182#[derive(Debug, Clone)]
183pub struct NetworkMetadata {
184    /// Creation timestamp
185    pub created_at: Instant,
186    /// Training history
187    pub training_history: Vec<TrainingEpoch>,
188    /// Performance metrics
189    pub performance_metrics: NetworkPerformanceMetrics,
190    /// Model version
191    pub version: String,
192}
193
194/// Training epoch information
195#[derive(Debug, Clone)]
196pub struct TrainingEpoch {
197    /// Epoch number
198    pub epoch: usize,
199    /// Training loss
200    pub training_loss: f64,
201    /// Validation loss
202    pub validation_loss: f64,
203    /// Learning rate
204    pub learning_rate: f64,
205    /// Training time
206    pub duration: Duration,
207    /// Additional metrics
208    pub metrics: HashMap<String, f64>,
209}
210
211/// Network performance metrics
212#[derive(Debug, Clone)]
213pub struct NetworkPerformanceMetrics {
214    /// Best validation loss achieved
215    pub best_validation_loss: f64,
216    /// Training convergence rate
217    pub convergence_rate: f64,
218    /// Generalization gap
219    pub generalization_gap: f64,
220    /// Parameter efficiency
221    pub parameter_efficiency: f64,
222}
223
224/// Network training state
225#[derive(Debug, Clone)]
226pub struct NetworkTrainingState {
227    /// Current epoch
228    pub current_epoch: usize,
229    /// Current learning rate
230    pub current_lr: f64,
231    /// Optimizer state
232    pub optimizer_state: OptimizerState,
233    /// Best model weights
234    pub best_weights: Option<Vec<Vec<Vec<f64>>>>,
235    /// Early stopping counter
236    pub early_stopping_counter: usize,
237}
238
239/// Optimizer state
240#[derive(Debug, Clone)]
241pub struct OptimizerState {
242    /// Momentum buffers (for SGD with momentum)
243    pub momentum_buffers: Vec<Vec<Vec<f64>>>,
244    /// First moment estimates (for Adam)
245    pub first_moments: Vec<Vec<Vec<f64>>>,
246    /// Second moment estimates (for Adam)
247    pub second_moments: Vec<Vec<Vec<f64>>>,
248    /// Iteration counter
249    pub iteration: usize,
250}
251
252impl EmbeddingDQN {
253    /// Create new DQN
254    pub fn new(layer_sizes: &[usize], seed: Option<u64>) -> RLEmbeddingResult<Self> {
255        let q_network = EmbeddingNetwork::new(layer_sizes, seed)?;
256        let target_network = q_network.clone();
257
258        let config = NetworkConfig {
259            layer_sizes: layer_sizes.to_vec(),
260            learning_rate: 0.001,
261            regularization: RegularizationConfig {
262                l1_strength: 0.0001,
263                l2_strength: 0.001,
264                dropout_rate: 0.1,
265                early_stopping_patience: 100,
266            },
267            optimizer: OptimizerType::Adam {
268                beta1: 0.9,
269                beta2: 0.999,
270            },
271            loss_function: LossFunction::MSE,
272        };
273
274        let training_state = NetworkTrainingState {
275            current_epoch: 0,
276            current_lr: 0.001,
277            optimizer_state: OptimizerState {
278                momentum_buffers: Vec::new(),
279                first_moments: Vec::new(),
280                second_moments: Vec::new(),
281                iteration: 0,
282            },
283            best_weights: None,
284            early_stopping_counter: 0,
285        };
286
287        Ok(Self {
288            q_network,
289            target_network,
290            config,
291            training_state,
292        })
293    }
294}
295
296impl EmbeddingPolicyNetwork {
297    /// Create new policy network
298    pub fn new(layer_sizes: &[usize], seed: Option<u64>) -> RLEmbeddingResult<Self> {
299        let actor_network = EmbeddingNetwork::new(layer_sizes, seed)?;
300        let critic_network = EmbeddingNetwork::new(layer_sizes, seed)?;
301
302        let config = NetworkConfig {
303            layer_sizes: layer_sizes.to_vec(),
304            learning_rate: 0.0001,
305            regularization: RegularizationConfig {
306                l1_strength: 0.0001,
307                l2_strength: 0.001,
308                dropout_rate: 0.1,
309                early_stopping_patience: 100,
310            },
311            optimizer: OptimizerType::Adam {
312                beta1: 0.9,
313                beta2: 0.999,
314            },
315            loss_function: LossFunction::MultiObjective,
316        };
317
318        let training_state = NetworkTrainingState {
319            current_epoch: 0,
320            current_lr: 0.0001,
321            optimizer_state: OptimizerState {
322                momentum_buffers: Vec::new(),
323                first_moments: Vec::new(),
324                second_moments: Vec::new(),
325                iteration: 0,
326            },
327            best_weights: None,
328            early_stopping_counter: 0,
329        };
330
331        Ok(Self {
332            actor_network,
333            critic_network,
334            config,
335            training_state,
336        })
337    }
338}
339
340impl EmbeddingNetwork {
341    /// Create new neural network
342    pub fn new(layer_sizes: &[usize], seed: Option<u64>) -> RLEmbeddingResult<Self> {
343        if layer_sizes.len() < 2 {
344            return Err(RLEmbeddingError::ConfigurationError(
345                "Network must have at least input and output layers".to_string(),
346            ));
347        }
348
349        let mut rng = match seed {
350            Some(s) => ChaCha8Rng::seed_from_u64(s),
351            None => ChaCha8Rng::seed_from_u64(thread_rng().random()),
352        };
353
354        let mut layers = Vec::new();
355
356        for i in 0..layer_sizes.len() - 1 {
357            let input_size = layer_sizes[i];
358            let output_size = layer_sizes[i + 1];
359
360            // Xavier initialization
361            let mut weights = vec![vec![0.0; input_size]; output_size];
362            let scale = (2.0 / input_size as f64).sqrt();
363
364            for row in &mut weights {
365                for weight in row {
366                    *weight = rng.random_range(-scale..scale);
367                }
368            }
369
370            let biases = vec![0.0; output_size];
371
372            let activation = if i == layer_sizes.len() - 2 {
373                ActivationFunction::Linear // Output layer
374            } else {
375                ActivationFunction::ReLU // Hidden layers
376            };
377
378            layers.push(NetworkLayer {
379                weights,
380                biases,
381                activation,
382                dropout_rate: 0.1,
383                batch_norm: None,
384            });
385        }
386
387        let input_size = layer_sizes[0];
388        let output_size = layer_sizes[layer_sizes.len() - 1];
389
390        let input_norm = NormalizationLayer {
391            means: vec![0.0; input_size],
392            stds: vec![1.0; input_size],
393            mins: vec![0.0; input_size],
394            maxs: vec![1.0; input_size],
395            norm_type: NormalizationType::StandardScore,
396        };
397
398        let output_scaling = NormalizationLayer {
399            means: vec![0.0; output_size],
400            stds: vec![1.0; output_size],
401            mins: vec![0.0; output_size],
402            maxs: vec![1.0; output_size],
403            norm_type: NormalizationType::None,
404        };
405
406        let metadata = NetworkMetadata {
407            created_at: Instant::now(),
408            training_history: Vec::new(),
409            performance_metrics: NetworkPerformanceMetrics {
410                best_validation_loss: f64::INFINITY,
411                convergence_rate: 0.0,
412                generalization_gap: 0.0,
413                parameter_efficiency: 0.0,
414            },
415            version: "1.0.0".to_string(),
416        };
417
418        Ok(Self {
419            layers,
420            input_norm,
421            output_scaling,
422            metadata,
423        })
424    }
425
426    /// Forward pass through network
427    pub fn forward(&self, input: &[f64]) -> RLEmbeddingResult<Vec<f64>> {
428        let mut activations = input.to_vec();
429
430        for layer in &self.layers {
431            activations = self.layer_forward(&activations, layer)?;
432        }
433
434        Ok(activations)
435    }
436
437    /// Forward pass through single layer
438    fn layer_forward(&self, input: &[f64], layer: &NetworkLayer) -> RLEmbeddingResult<Vec<f64>> {
439        if input.len() != layer.weights[0].len() {
440            return Err(RLEmbeddingError::NeuralNetworkError(format!(
441                "Input size {} doesn't match layer input size {}",
442                input.len(),
443                layer.weights[0].len()
444            )));
445        }
446
447        let mut output = Vec::new();
448
449        for (neuron_weights, &bias) in layer.weights.iter().zip(&layer.biases) {
450            let mut activation = bias;
451
452            for (&inp, &weight) in input.iter().zip(neuron_weights) {
453                activation += inp * weight;
454            }
455
456            // Apply activation function
457            activation = match layer.activation {
458                ActivationFunction::ReLU => activation.max(0.0),
459                ActivationFunction::LeakyReLU(alpha) => {
460                    if activation > 0.0 {
461                        activation
462                    } else {
463                        alpha * activation
464                    }
465                }
466                ActivationFunction::Tanh => activation.tanh(),
467                ActivationFunction::Sigmoid => 1.0 / (1.0 + (-activation).exp()),
468                ActivationFunction::Swish => activation / (1.0 + (-activation).exp()),
469                ActivationFunction::Linear => activation,
470            };
471
472            output.push(activation);
473        }
474
475        Ok(output)
476    }
477}