1use 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#[derive(Debug, Clone)]
13pub struct EmbeddingDQN {
14 pub q_network: EmbeddingNetwork,
16 pub target_network: EmbeddingNetwork,
18 pub config: NetworkConfig,
20 pub training_state: NetworkTrainingState,
22}
23
24#[derive(Debug, Clone)]
26pub struct EmbeddingPolicyNetwork {
27 pub actor_network: EmbeddingNetwork,
29 pub critic_network: EmbeddingNetwork,
31 pub config: NetworkConfig,
33 pub training_state: NetworkTrainingState,
35}
36
37#[derive(Debug, Clone)]
39pub struct EmbeddingNetwork {
40 pub layers: Vec<NetworkLayer>,
42 pub input_norm: NormalizationLayer,
44 pub output_scaling: NormalizationLayer,
46 pub metadata: NetworkMetadata,
48}
49
50#[derive(Debug, Clone)]
52pub struct NetworkLayer {
53 pub weights: Vec<Vec<f64>>,
55 pub biases: Vec<f64>,
57 pub activation: ActivationFunction,
59 pub dropout_rate: f64,
61 pub batch_norm: Option<BatchNormalization>,
63}
64
65#[derive(Debug, Clone, PartialEq)]
67pub enum ActivationFunction {
68 ReLU,
70 LeakyReLU(f64),
72 Tanh,
74 Sigmoid,
76 Swish,
78 Linear,
80}
81
82#[derive(Debug, Clone)]
84pub struct BatchNormalization {
85 pub running_mean: Vec<f64>,
87 pub running_var: Vec<f64>,
89 pub gamma: Vec<f64>,
91 pub beta: Vec<f64>,
93 pub epsilon: f64,
95 pub momentum: f64,
97}
98
99#[derive(Debug, Clone)]
101pub struct NormalizationLayer {
102 pub means: Vec<f64>,
104 pub stds: Vec<f64>,
106 pub mins: Vec<f64>,
108 pub maxs: Vec<f64>,
110 pub norm_type: NormalizationType,
112}
113
114#[derive(Debug, Clone, PartialEq, Eq)]
116pub enum NormalizationType {
117 StandardScore,
119 MinMax,
121 Robust,
123 None,
125}
126
127#[derive(Debug, Clone)]
129pub struct NetworkConfig {
130 pub layer_sizes: Vec<usize>,
132 pub learning_rate: f64,
134 pub regularization: RegularizationConfig,
136 pub optimizer: OptimizerType,
138 pub loss_function: LossFunction,
140}
141
142#[derive(Debug, Clone)]
144pub struct RegularizationConfig {
145 pub l1_strength: f64,
147 pub l2_strength: f64,
149 pub dropout_rate: f64,
151 pub early_stopping_patience: usize,
153}
154
155#[derive(Debug, Clone, PartialEq)]
157pub enum OptimizerType {
158 SGD,
160 Adam { beta1: f64, beta2: f64 },
162 RMSprop { decay_rate: f64 },
164 AdaGrad,
166}
167
168#[derive(Debug, Clone, PartialEq)]
170pub enum LossFunction {
171 MSE,
173 Huber { delta: f64 },
175 CrossEntropy,
177 MultiObjective,
179}
180
181#[derive(Debug, Clone)]
183pub struct NetworkMetadata {
184 pub created_at: Instant,
186 pub training_history: Vec<TrainingEpoch>,
188 pub performance_metrics: NetworkPerformanceMetrics,
190 pub version: String,
192}
193
194#[derive(Debug, Clone)]
196pub struct TrainingEpoch {
197 pub epoch: usize,
199 pub training_loss: f64,
201 pub validation_loss: f64,
203 pub learning_rate: f64,
205 pub duration: Duration,
207 pub metrics: HashMap<String, f64>,
209}
210
211#[derive(Debug, Clone)]
213pub struct NetworkPerformanceMetrics {
214 pub best_validation_loss: f64,
216 pub convergence_rate: f64,
218 pub generalization_gap: f64,
220 pub parameter_efficiency: f64,
222}
223
224#[derive(Debug, Clone)]
226pub struct NetworkTrainingState {
227 pub current_epoch: usize,
229 pub current_lr: f64,
231 pub optimizer_state: OptimizerState,
233 pub best_weights: Option<Vec<Vec<Vec<f64>>>>,
235 pub early_stopping_counter: usize,
237}
238
239#[derive(Debug, Clone)]
241pub struct OptimizerState {
242 pub momentum_buffers: Vec<Vec<Vec<f64>>>,
244 pub first_moments: Vec<Vec<Vec<f64>>>,
246 pub second_moments: Vec<Vec<Vec<f64>>>,
248 pub iteration: usize,
250}
251
252impl EmbeddingDQN {
253 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 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 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 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 } else {
375 ActivationFunction::ReLU };
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 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 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 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}