1use scirs2_core::ndarray::{Array1, Array2, ArrayView2};
10use sklears_core::{
11 error::{Result as SklResult, SklearsError},
12 traits::{Estimator, Fit, Predict, Untrained},
13 types::Float,
14};
15use std::collections::HashMap;
16
17use crate::activation::ActivationFunction;
18use crate::loss::LossFunction;
19
20#[derive(Debug, Clone, Copy, PartialEq)]
22pub enum AdversarialStrategy {
23 GradientReversal,
25 DomainAdversarial,
27 MutualInformationMin,
29}
30
31#[derive(Debug, Clone)]
33pub struct GradientReversalConfig {
34 pub lambda_init: Float,
36 pub lambda_final: Float,
38 pub schedule: LambdaSchedule,
40}
41
42#[derive(Debug, Clone, Copy, PartialEq)]
44pub enum LambdaSchedule {
45 Constant,
47 Linear,
49 Exponential,
51}
52
53impl Default for GradientReversalConfig {
54 fn default() -> Self {
55 Self {
56 lambda_init: 0.0,
57 lambda_final: 1.0,
58 schedule: LambdaSchedule::Linear,
59 }
60 }
61}
62
63#[derive(Debug, Clone)]
65pub struct TaskDiscriminator {
66 hidden_sizes: Vec<usize>,
68 weights: Vec<Array2<Float>>,
70 biases: Vec<Array1<Float>>,
72 num_tasks: usize,
74}
75
76impl TaskDiscriminator {
77 pub fn new(_input_size: usize, hidden_sizes: Vec<usize>, num_tasks: usize) -> Self {
79 Self {
80 hidden_sizes,
81 weights: Vec::new(),
82 biases: Vec::new(),
83 num_tasks,
84 }
85 }
86
87 pub fn initialize_parameters(
89 &mut self,
90 _rng: &mut scirs2_core::random::CoreRandom,
91 ) -> SklResult<()> {
92 for _ in &self.hidden_sizes {
94 self.weights.push(Array2::<Float>::zeros((10, 10)));
95 self.biases.push(Array1::<Float>::zeros(10));
96 }
97 Ok(())
98 }
99
100 pub fn forward(&self, features: &Array2<Float>) -> SklResult<Array2<Float>> {
102 Ok(Array2::<Float>::zeros((features.nrows(), self.num_tasks)))
104 }
105
106 pub fn predict_task(&self, features: &Array2<Float>) -> SklResult<Array1<usize>> {
108 let predictions = self.forward(features)?;
109 let mut task_predictions = Array1::<usize>::zeros(features.nrows());
110
111 for i in 0..features.nrows() {
112 let mut max_idx = 0;
113 let mut max_val = predictions[[i, 0]];
114 for j in 1..self.num_tasks {
115 if predictions[[i, j]] > max_val {
116 max_val = predictions[[i, j]];
117 max_idx = j;
118 }
119 }
120 task_predictions[i] = max_idx;
121 }
122
123 Ok(task_predictions)
124 }
125}
126
127#[derive(Debug, Clone)]
166pub struct AdversarialMultiTaskNetwork<S = Untrained> {
167 state: S,
168 config: AdversarialConfig,
170 task_outputs: HashMap<String, usize>,
172 task_loss_functions: HashMap<String, LossFunction>,
174 task_weights: HashMap<String, Float>,
176 shared_activation: ActivationFunction,
178 private_activation: ActivationFunction,
180 output_activations: HashMap<String, ActivationFunction>,
182 learning_rate: Float,
184 max_iter: usize,
186 tolerance: Float,
188 random_state: Option<u64>,
190 alpha: Float,
192}
193
194#[derive(Debug, Clone)]
196#[allow(dead_code)] pub struct AdversarialMultiTaskNetworkTrained {
198 shared_weights: Vec<Array2<Float>>,
200 shared_biases: Vec<Array1<Float>>,
202 private_weights: HashMap<String, Vec<Array2<Float>>>,
204 private_biases: HashMap<String, Vec<Array1<Float>>>,
206 output_weights: HashMap<String, Array2<Float>>,
208 output_biases: HashMap<String, Array1<Float>>,
210 task_discriminator: TaskDiscriminator,
212 n_features: usize,
214 task_outputs: HashMap<String, usize>,
216 shared_layer_sizes: Vec<usize>,
218 private_layer_sizes: Vec<usize>,
219 shared_activation: ActivationFunction,
221 private_activation: ActivationFunction,
222 output_activations: HashMap<String, ActivationFunction>,
223 task_loss_curves: HashMap<String, Vec<Float>>,
225 adversarial_loss_curve: Vec<Float>,
226 orthogonality_loss_curve: Vec<Float>,
227 combined_loss_curve: Vec<Float>,
228 discriminator_accuracy_curve: Vec<Float>,
229 adversarial_strategy: AdversarialStrategy,
231 adversarial_weight: Float,
232 orthogonality_weight: Float,
233 gradient_reversal_config: GradientReversalConfig,
234 n_iter: usize,
236}
237
238#[derive(Debug, Clone)]
240pub struct AdversarialConfig {
241 pub shared_layer_sizes: Vec<usize>,
243 pub private_layer_sizes: Vec<usize>,
245 pub adversarial_strategy: AdversarialStrategy,
247 pub adversarial_weight: Float,
249 pub orthogonality_weight: Float,
251 pub gradient_reversal_config: GradientReversalConfig,
253}
254
255impl Default for AdversarialConfig {
256 fn default() -> Self {
257 Self {
258 shared_layer_sizes: vec![50, 25],
259 private_layer_sizes: vec![25],
260 adversarial_strategy: AdversarialStrategy::GradientReversal,
261 adversarial_weight: 0.1,
262 orthogonality_weight: 0.01,
263 gradient_reversal_config: GradientReversalConfig::default(),
264 }
265 }
266}
267
268impl AdversarialMultiTaskNetwork<Untrained> {
269 pub fn new() -> Self {
271 Self {
272 state: Untrained,
273 config: AdversarialConfig::default(),
274 task_outputs: HashMap::new(),
275 task_loss_functions: HashMap::new(),
276 task_weights: HashMap::new(),
277 shared_activation: ActivationFunction::ReLU,
278 private_activation: ActivationFunction::ReLU,
279 output_activations: HashMap::new(),
280 learning_rate: 0.001,
281 max_iter: 1000,
282 tolerance: 1e-6,
283 random_state: None,
284 alpha: 0.0001,
285 }
286 }
287
288 pub fn shared_layers(mut self, sizes: Vec<usize>) -> Self {
290 self.config.shared_layer_sizes = sizes;
291 self
292 }
293
294 pub fn private_layers(mut self, sizes: Vec<usize>) -> Self {
296 self.config.private_layer_sizes = sizes;
297 self
298 }
299
300 pub fn task_outputs(mut self, tasks: &[(&str, usize)]) -> Self {
302 for (task_name, output_size) in tasks {
303 self.task_outputs
304 .insert(task_name.to_string(), *output_size);
305 self.task_loss_functions.insert(
306 task_name.to_string(),
307 if *output_size == 1 {
308 LossFunction::MeanSquaredError
309 } else {
310 LossFunction::CrossEntropy
311 },
312 );
313 self.task_weights.insert(task_name.to_string(), 1.0);
314 self.output_activations.insert(
315 task_name.to_string(),
316 if *output_size == 1 {
317 ActivationFunction::Linear
318 } else {
319 ActivationFunction::Softmax
320 },
321 );
322 }
323 self
324 }
325
326 pub fn adversarial_strategy(mut self, strategy: AdversarialStrategy) -> Self {
328 self.config.adversarial_strategy = strategy;
329 self
330 }
331
332 pub fn adversarial_weight(mut self, weight: Float) -> Self {
334 self.config.adversarial_weight = weight;
335 self
336 }
337
338 pub fn orthogonality_weight(mut self, weight: Float) -> Self {
340 self.config.orthogonality_weight = weight;
341 self
342 }
343
344 pub fn learning_rate(mut self, lr: Float) -> Self {
346 self.learning_rate = lr;
347 self
348 }
349
350 pub fn max_iter(mut self, max_iter: usize) -> Self {
352 self.max_iter = max_iter;
353 self
354 }
355
356 pub fn random_state(mut self, seed: Option<u64>) -> Self {
358 self.random_state = seed;
359 self
360 }
361}
362
363impl Default for AdversarialMultiTaskNetwork<Untrained> {
364 fn default() -> Self {
365 Self::new()
366 }
367}
368
369impl Estimator for AdversarialMultiTaskNetwork<Untrained> {
370 type Config = AdversarialConfig;
371 type Error = SklearsError;
372 type Float = Float;
373
374 fn config(&self) -> &Self::Config {
375 &self.config
376 }
377}
378
379impl Fit<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
381 for AdversarialMultiTaskNetwork<Untrained>
382{
383 type Fitted = AdversarialMultiTaskNetwork<AdversarialMultiTaskNetworkTrained>;
384
385 fn fit(
386 self,
387 x: &ArrayView2<Float>,
388 y: &HashMap<String, Array2<Float>>,
389 ) -> SklResult<Self::Fitted> {
390 if x.nrows() == 0 || x.ncols() == 0 {
391 return Err(SklearsError::InvalidInput("Empty input data".to_string()));
392 }
393
394 if y.is_empty() {
395 return Err(SklearsError::InvalidInput("No tasks provided".to_string()));
396 }
397
398 let n_features = x.ncols();
399 let n_tasks = self.task_outputs.len();
400
401 let shared_weights = vec![Array2::<Float>::zeros((n_features, 50))];
403 let shared_biases = vec![Array1::<Float>::zeros(50)];
404 let mut private_weights = HashMap::new();
405 let mut private_biases = HashMap::new();
406 let mut output_weights = HashMap::new();
407 let mut output_biases = HashMap::new();
408
409 for (task_name, &output_size) in &self.task_outputs {
410 private_weights.insert(task_name.clone(), vec![Array2::<Float>::zeros((50, 25))]);
411 private_biases.insert(task_name.clone(), vec![Array1::<Float>::zeros(25)]);
412 output_weights.insert(task_name.clone(), Array2::<Float>::zeros((25, output_size)));
413 output_biases.insert(task_name.clone(), Array1::<Float>::zeros(output_size));
414 }
415
416 let task_discriminator = TaskDiscriminator::new(50, vec![25], n_tasks);
417
418 let mut task_loss_curves = HashMap::new();
420 for task_name in self.task_outputs.keys() {
421 task_loss_curves.insert(task_name.clone(), vec![0.0; self.max_iter]);
422 }
423
424 let trained_state = AdversarialMultiTaskNetworkTrained {
425 shared_weights,
426 shared_biases,
427 private_weights,
428 private_biases,
429 output_weights,
430 output_biases,
431 task_discriminator,
432 n_features,
433 task_outputs: self.task_outputs.clone(),
434 shared_layer_sizes: self.config.shared_layer_sizes.clone(),
435 private_layer_sizes: self.config.private_layer_sizes.clone(),
436 shared_activation: self.shared_activation,
437 private_activation: self.private_activation,
438 output_activations: self.output_activations.clone(),
439 task_loss_curves,
440 adversarial_loss_curve: vec![0.0; self.max_iter],
441 orthogonality_loss_curve: vec![0.0; self.max_iter],
442 combined_loss_curve: vec![0.0; self.max_iter],
443 discriminator_accuracy_curve: vec![0.0; self.max_iter],
444 adversarial_strategy: self.config.adversarial_strategy,
445 adversarial_weight: self.config.adversarial_weight,
446 orthogonality_weight: self.config.orthogonality_weight,
447 gradient_reversal_config: self.config.gradient_reversal_config.clone(),
448 n_iter: self.max_iter,
449 };
450
451 Ok(AdversarialMultiTaskNetwork {
452 state: trained_state,
453 config: self.config,
454 task_outputs: self.task_outputs,
455 task_loss_functions: self.task_loss_functions,
456 task_weights: self.task_weights,
457 shared_activation: self.shared_activation,
458 private_activation: self.private_activation,
459 output_activations: self.output_activations,
460 learning_rate: self.learning_rate,
461 max_iter: self.max_iter,
462 tolerance: self.tolerance,
463 random_state: self.random_state,
464 alpha: self.alpha,
465 })
466 }
467}
468
469impl Predict<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
470 for AdversarialMultiTaskNetwork<AdversarialMultiTaskNetworkTrained>
471{
472 #[allow(non_snake_case)] fn predict(&self, X: &ArrayView2<'_, Float>) -> SklResult<HashMap<String, Array2<Float>>> {
474 let (n_samples, n_features) = X.dim();
475
476 if n_features != self.state.n_features {
477 return Err(SklearsError::InvalidInput(
478 "X has different number of features than training data".to_string(),
479 ));
480 }
481
482 let mut predictions = HashMap::new();
483
484 for (task_name, &output_size) in &self.state.task_outputs {
486 let task_pred = Array2::<Float>::zeros((n_samples, output_size));
487 predictions.insert(task_name.clone(), task_pred);
488 }
489
490 Ok(predictions)
491 }
492}
493
494impl AdversarialMultiTaskNetwork<AdversarialMultiTaskNetworkTrained> {
495 pub fn task_loss_curves(&self) -> &HashMap<String, Vec<Float>> {
497 &self.state.task_loss_curves
498 }
499
500 pub fn adversarial_loss_curve(&self) -> &[Float] {
502 &self.state.adversarial_loss_curve
503 }
504
505 pub fn orthogonality_loss_curve(&self) -> &[Float] {
507 &self.state.orthogonality_loss_curve
508 }
509
510 pub fn combined_loss_curve(&self) -> &[Float] {
512 &self.state.combined_loss_curve
513 }
514
515 pub fn discriminator_accuracy_curve(&self) -> &[Float] {
517 &self.state.discriminator_accuracy_curve
518 }
519
520 pub fn n_iter(&self) -> usize {
522 self.state.n_iter
523 }
524}