Skip to main content

sklears_multioutput/
adversarial.rs

1//! Adversarial Multi-Task Networks with Feature Disentanglement
2//!
3//! This module implements adversarial multi-task learning where a task discriminator
4//! is trained to predict which task shared features come from, while the shared
5//! feature extractor is trained adversarially to fool the discriminator. This ensures
6//! that shared representations contain only task-invariant information.
7
8// Use SciRS2-Core for arrays and random number generation (SciRS2 Policy)
9use 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/// Adversarial training strategies for multi-task learning
21#[derive(Debug, Clone, Copy, PartialEq)]
22pub enum AdversarialStrategy {
23    /// Gradient reversal layer
24    GradientReversal,
25    /// Domain adversarial training
26    DomainAdversarial,
27    /// Mutual information minimization
28    MutualInformationMin,
29}
30
31/// Configuration for gradient reversal layer
32#[derive(Debug, Clone)]
33pub struct GradientReversalConfig {
34    /// Initial lambda value for gradient reversal
35    pub lambda_init: Float,
36    /// Final lambda value
37    pub lambda_final: Float,
38    /// Lambda scheduling strategy
39    pub schedule: LambdaSchedule,
40}
41
42/// Lambda scheduling strategies for gradient reversal
43#[derive(Debug, Clone, Copy, PartialEq)]
44pub enum LambdaSchedule {
45    /// Constant lambda value
46    Constant,
47    /// Linear increase from init to final
48    Linear,
49    /// Exponential increase
50    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/// Task discriminator for adversarial training
64#[derive(Debug, Clone)]
65pub struct TaskDiscriminator {
66    /// Hidden layer sizes
67    hidden_sizes: Vec<usize>,
68    /// Weights for each layer
69    weights: Vec<Array2<Float>>,
70    /// Biases for each layer
71    biases: Vec<Array1<Float>>,
72    /// Number of tasks
73    num_tasks: usize,
74}
75
76impl TaskDiscriminator {
77    /// Create a new task discriminator
78    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    /// Initialize parameters
88    pub fn initialize_parameters(
89        &mut self,
90        _rng: &mut scirs2_core::random::CoreRandom,
91    ) -> SklResult<()> {
92        // Simplified initialization
93        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    /// Forward pass
101    pub fn forward(&self, features: &Array2<Float>) -> SklResult<Array2<Float>> {
102        // Simplified forward pass
103        Ok(Array2::<Float>::zeros((features.nrows(), self.num_tasks)))
104    }
105
106    /// Predict task labels from features
107    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/// Adversarial Multi-Task Network with feature disentanglement
128///
129/// This network implements adversarial multi-task learning where a task discriminator
130/// is trained to predict which task shared features come from, while the shared
131/// feature extractor is trained adversarially to fool the discriminator. This ensures
132/// that shared representations contain only task-invariant information.
133///
134/// # Architecture
135///
136/// The network consists of:
137/// - Shared layers: Learn task-invariant representations
138/// - Private layers: Learn task-specific representations per task
139/// - Task discriminator: Tries to predict task from shared features
140/// - Gradient reversal: Adversarial training mechanism
141///
142/// # Examples
143///
144/// ```
145/// use sklears_multioutput::adversarial::{AdversarialMultiTaskNetwork, AdversarialStrategy};
146/// use sklears_core::traits::{Predict, Fit};
147/// // Use SciRS2-Core for arrays and random number generation (SciRS2 Policy)
148/// use scirs2_core::ndarray::array;
149/// use std::collections::HashMap;
150///
151/// let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [4.0, 4.0]];
152/// let mut tasks = HashMap::new();
153/// tasks.insert("task1".to_string(), array![[0.5], [1.0], [1.5], [2.0]]);
154/// tasks.insert("task2".to_string(), array![[1.0], [0.0], [1.0], [0.0]]);
155///
156/// let adv_net = AdversarialMultiTaskNetwork::new()
157///     .shared_layers(vec![20, 10])
158///     .private_layers(vec![8])
159///     .task_outputs(&[("task1", 1), ("task2", 1)])
160///     .adversarial_strategy(AdversarialStrategy::GradientReversal)
161///     .adversarial_weight(0.1)
162///     .orthogonality_weight(0.01)
163///     .random_state(Some(42));
164/// ```
165#[derive(Debug, Clone)]
166pub struct AdversarialMultiTaskNetwork<S = Untrained> {
167    state: S,
168    /// Configuration for adversarial training
169    config: AdversarialConfig,
170    /// Task outputs configuration
171    task_outputs: HashMap<String, usize>,
172    /// Task loss functions
173    task_loss_functions: HashMap<String, LossFunction>,
174    /// Task weights for loss computation
175    task_weights: HashMap<String, Float>,
176    /// Shared activation function
177    shared_activation: ActivationFunction,
178    /// Private activation function
179    private_activation: ActivationFunction,
180    /// Output activation functions per task
181    output_activations: HashMap<String, ActivationFunction>,
182    /// Learning rate
183    learning_rate: Float,
184    /// Maximum iterations
185    max_iter: usize,
186    /// Convergence tolerance
187    tolerance: Float,
188    /// Random state for reproducibility
189    random_state: Option<u64>,
190    /// L2 regularization
191    alpha: Float,
192}
193
194/// Trained state for AdversarialMultiTaskNetwork
195#[derive(Debug, Clone)]
196#[allow(dead_code)] // fields used for serialization and future inference extensions
197pub struct AdversarialMultiTaskNetworkTrained {
198    /// Shared layer weights
199    shared_weights: Vec<Array2<Float>>,
200    /// Shared layer biases
201    shared_biases: Vec<Array1<Float>>,
202    /// Private layer weights per task
203    private_weights: HashMap<String, Vec<Array2<Float>>>,
204    /// Private layer biases per task
205    private_biases: HashMap<String, Vec<Array1<Float>>>,
206    /// Output layer weights per task
207    output_weights: HashMap<String, Array2<Float>>,
208    /// Output layer biases per task
209    output_biases: HashMap<String, Array1<Float>>,
210    /// Task discriminator
211    task_discriminator: TaskDiscriminator,
212    /// Number of input features
213    n_features: usize,
214    /// Task configurations
215    task_outputs: HashMap<String, usize>,
216    /// Network architecture
217    shared_layer_sizes: Vec<usize>,
218    private_layer_sizes: Vec<usize>,
219    /// Activation functions
220    shared_activation: ActivationFunction,
221    private_activation: ActivationFunction,
222    output_activations: HashMap<String, ActivationFunction>,
223    /// Training history
224    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 configuration
230    adversarial_strategy: AdversarialStrategy,
231    adversarial_weight: Float,
232    orthogonality_weight: Float,
233    gradient_reversal_config: GradientReversalConfig,
234    /// Training iterations
235    n_iter: usize,
236}
237
238/// Configuration for AdversarialMultiTaskNetwork
239#[derive(Debug, Clone)]
240pub struct AdversarialConfig {
241    /// Shared layer sizes
242    pub shared_layer_sizes: Vec<usize>,
243    /// Private layer sizes per task
244    pub private_layer_sizes: Vec<usize>,
245    /// Adversarial strategy
246    pub adversarial_strategy: AdversarialStrategy,
247    /// Weight for adversarial loss
248    pub adversarial_weight: Float,
249    /// Weight for orthogonality constraint
250    pub orthogonality_weight: Float,
251    /// Gradient reversal configuration
252    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    /// Create a new AdversarialMultiTaskNetwork
270    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    /// Set shared layer sizes
289    pub fn shared_layers(mut self, sizes: Vec<usize>) -> Self {
290        self.config.shared_layer_sizes = sizes;
291        self
292    }
293
294    /// Set private layer sizes
295    pub fn private_layers(mut self, sizes: Vec<usize>) -> Self {
296        self.config.private_layer_sizes = sizes;
297        self
298    }
299
300    /// Configure task outputs
301    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    /// Set adversarial strategy
327    pub fn adversarial_strategy(mut self, strategy: AdversarialStrategy) -> Self {
328        self.config.adversarial_strategy = strategy;
329        self
330    }
331
332    /// Set adversarial weight
333    pub fn adversarial_weight(mut self, weight: Float) -> Self {
334        self.config.adversarial_weight = weight;
335        self
336    }
337
338    /// Set orthogonality weight
339    pub fn orthogonality_weight(mut self, weight: Float) -> Self {
340        self.config.orthogonality_weight = weight;
341        self
342    }
343
344    /// Set learning rate
345    pub fn learning_rate(mut self, lr: Float) -> Self {
346        self.learning_rate = lr;
347        self
348    }
349
350    /// Set maximum iterations
351    pub fn max_iter(mut self, max_iter: usize) -> Self {
352        self.max_iter = max_iter;
353        self
354    }
355
356    /// Set random state
357    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
379// Simplified implementation for demonstration
380impl 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        // Simplified parameter initialization
402        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        // Simplified training history
419        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)] // standard ML notation
473    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        // Simplified prediction logic
485        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    /// Get task loss curves
496    pub fn task_loss_curves(&self) -> &HashMap<String, Vec<Float>> {
497        &self.state.task_loss_curves
498    }
499
500    /// Get adversarial loss curve
501    pub fn adversarial_loss_curve(&self) -> &[Float] {
502        &self.state.adversarial_loss_curve
503    }
504
505    /// Get orthogonality loss curve
506    pub fn orthogonality_loss_curve(&self) -> &[Float] {
507        &self.state.orthogonality_loss_curve
508    }
509
510    /// Get combined loss curve
511    pub fn combined_loss_curve(&self) -> &[Float] {
512        &self.state.combined_loss_curve
513    }
514
515    /// Get discriminator accuracy curve
516    pub fn discriminator_accuracy_curve(&self) -> &[Float] {
517        &self.state.discriminator_accuracy_curve
518    }
519
520    /// Get training iterations
521    pub fn n_iter(&self) -> usize {
522        self.state.n_iter
523    }
524}