use scirs2_core::ndarray::{Array1, Array2, ArrayView2};
use sklears_core::{
error::{Result as SklResult, SklearsError},
traits::{Estimator, Fit, Predict, Untrained},
types::Float,
};
use std::collections::HashMap;
use crate::activation::ActivationFunction;
use crate::loss::LossFunction;
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum AdversarialStrategy {
GradientReversal,
DomainAdversarial,
MutualInformationMin,
}
#[derive(Debug, Clone)]
pub struct GradientReversalConfig {
pub lambda_init: Float,
pub lambda_final: Float,
pub schedule: LambdaSchedule,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub enum LambdaSchedule {
Constant,
Linear,
Exponential,
}
impl Default for GradientReversalConfig {
fn default() -> Self {
Self {
lambda_init: 0.0,
lambda_final: 1.0,
schedule: LambdaSchedule::Linear,
}
}
}
#[derive(Debug, Clone)]
pub struct TaskDiscriminator {
hidden_sizes: Vec<usize>,
weights: Vec<Array2<Float>>,
biases: Vec<Array1<Float>>,
num_tasks: usize,
}
impl TaskDiscriminator {
pub fn new(input_size: usize, hidden_sizes: Vec<usize>, num_tasks: usize) -> Self {
Self {
hidden_sizes,
weights: Vec::new(),
biases: Vec::new(),
num_tasks,
}
}
pub fn initialize_parameters(
&mut self,
rng: &mut scirs2_core::random::CoreRandom,
) -> SklResult<()> {
for _ in &self.hidden_sizes {
self.weights.push(Array2::<Float>::zeros((10, 10)));
self.biases.push(Array1::<Float>::zeros(10));
}
Ok(())
}
pub fn forward(&self, features: &Array2<Float>) -> SklResult<Array2<Float>> {
Ok(Array2::<Float>::zeros((features.nrows(), self.num_tasks)))
}
pub fn predict_task(&self, features: &Array2<Float>) -> SklResult<Array1<usize>> {
let predictions = self.forward(features)?;
let mut task_predictions = Array1::<usize>::zeros(features.nrows());
for i in 0..features.nrows() {
let mut max_idx = 0;
let mut max_val = predictions[[i, 0]];
for j in 1..self.num_tasks {
if predictions[[i, j]] > max_val {
max_val = predictions[[i, j]];
max_idx = j;
}
}
task_predictions[i] = max_idx;
}
Ok(task_predictions)
}
}
#[derive(Debug, Clone)]
pub struct AdversarialMultiTaskNetwork<S = Untrained> {
state: S,
config: AdversarialConfig,
task_outputs: HashMap<String, usize>,
task_loss_functions: HashMap<String, LossFunction>,
task_weights: HashMap<String, Float>,
shared_activation: ActivationFunction,
private_activation: ActivationFunction,
output_activations: HashMap<String, ActivationFunction>,
learning_rate: Float,
max_iter: usize,
tolerance: Float,
random_state: Option<u64>,
alpha: Float,
}
#[derive(Debug, Clone)]
pub struct AdversarialMultiTaskNetworkTrained {
shared_weights: Vec<Array2<Float>>,
shared_biases: Vec<Array1<Float>>,
private_weights: HashMap<String, Vec<Array2<Float>>>,
private_biases: HashMap<String, Vec<Array1<Float>>>,
output_weights: HashMap<String, Array2<Float>>,
output_biases: HashMap<String, Array1<Float>>,
task_discriminator: TaskDiscriminator,
n_features: usize,
task_outputs: HashMap<String, usize>,
shared_layer_sizes: Vec<usize>,
private_layer_sizes: Vec<usize>,
shared_activation: ActivationFunction,
private_activation: ActivationFunction,
output_activations: HashMap<String, ActivationFunction>,
task_loss_curves: HashMap<String, Vec<Float>>,
adversarial_loss_curve: Vec<Float>,
orthogonality_loss_curve: Vec<Float>,
combined_loss_curve: Vec<Float>,
discriminator_accuracy_curve: Vec<Float>,
adversarial_strategy: AdversarialStrategy,
adversarial_weight: Float,
orthogonality_weight: Float,
gradient_reversal_config: GradientReversalConfig,
n_iter: usize,
}
#[derive(Debug, Clone)]
pub struct AdversarialConfig {
pub shared_layer_sizes: Vec<usize>,
pub private_layer_sizes: Vec<usize>,
pub adversarial_strategy: AdversarialStrategy,
pub adversarial_weight: Float,
pub orthogonality_weight: Float,
pub gradient_reversal_config: GradientReversalConfig,
}
impl Default for AdversarialConfig {
fn default() -> Self {
Self {
shared_layer_sizes: vec![50, 25],
private_layer_sizes: vec![25],
adversarial_strategy: AdversarialStrategy::GradientReversal,
adversarial_weight: 0.1,
orthogonality_weight: 0.01,
gradient_reversal_config: GradientReversalConfig::default(),
}
}
}
impl AdversarialMultiTaskNetwork<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
config: AdversarialConfig::default(),
task_outputs: HashMap::new(),
task_loss_functions: HashMap::new(),
task_weights: HashMap::new(),
shared_activation: ActivationFunction::ReLU,
private_activation: ActivationFunction::ReLU,
output_activations: HashMap::new(),
learning_rate: 0.001,
max_iter: 1000,
tolerance: 1e-6,
random_state: None,
alpha: 0.0001,
}
}
pub fn shared_layers(mut self, sizes: Vec<usize>) -> Self {
self.config.shared_layer_sizes = sizes;
self
}
pub fn private_layers(mut self, sizes: Vec<usize>) -> Self {
self.config.private_layer_sizes = sizes;
self
}
pub fn task_outputs(mut self, tasks: &[(&str, usize)]) -> Self {
for (task_name, output_size) in tasks {
self.task_outputs
.insert(task_name.to_string(), *output_size);
self.task_loss_functions.insert(
task_name.to_string(),
if *output_size == 1 {
LossFunction::MeanSquaredError
} else {
LossFunction::CrossEntropy
},
);
self.task_weights.insert(task_name.to_string(), 1.0);
self.output_activations.insert(
task_name.to_string(),
if *output_size == 1 {
ActivationFunction::Linear
} else {
ActivationFunction::Softmax
},
);
}
self
}
pub fn adversarial_strategy(mut self, strategy: AdversarialStrategy) -> Self {
self.config.adversarial_strategy = strategy;
self
}
pub fn adversarial_weight(mut self, weight: Float) -> Self {
self.config.adversarial_weight = weight;
self
}
pub fn orthogonality_weight(mut self, weight: Float) -> Self {
self.config.orthogonality_weight = weight;
self
}
pub fn learning_rate(mut self, lr: Float) -> Self {
self.learning_rate = lr;
self
}
pub fn max_iter(mut self, max_iter: usize) -> Self {
self.max_iter = max_iter;
self
}
pub fn random_state(mut self, seed: Option<u64>) -> Self {
self.random_state = seed;
self
}
}
impl Default for AdversarialMultiTaskNetwork<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl Estimator for AdversarialMultiTaskNetwork<Untrained> {
type Config = AdversarialConfig;
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&self.config
}
}
impl Fit<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
for AdversarialMultiTaskNetwork<Untrained>
{
type Fitted = AdversarialMultiTaskNetwork<AdversarialMultiTaskNetworkTrained>;
fn fit(
self,
x: &ArrayView2<Float>,
y: &HashMap<String, Array2<Float>>,
) -> SklResult<Self::Fitted> {
if x.nrows() == 0 || x.ncols() == 0 {
return Err(SklearsError::InvalidInput("Empty input data".to_string()));
}
if y.is_empty() {
return Err(SklearsError::InvalidInput("No tasks provided".to_string()));
}
let n_features = x.ncols();
let n_tasks = self.task_outputs.len();
let shared_weights = vec![Array2::<Float>::zeros((n_features, 50))];
let shared_biases = vec![Array1::<Float>::zeros(50)];
let mut private_weights = HashMap::new();
let mut private_biases = HashMap::new();
let mut output_weights = HashMap::new();
let mut output_biases = HashMap::new();
for (task_name, &output_size) in &self.task_outputs {
private_weights.insert(task_name.clone(), vec![Array2::<Float>::zeros((50, 25))]);
private_biases.insert(task_name.clone(), vec![Array1::<Float>::zeros(25)]);
output_weights.insert(task_name.clone(), Array2::<Float>::zeros((25, output_size)));
output_biases.insert(task_name.clone(), Array1::<Float>::zeros(output_size));
}
let task_discriminator = TaskDiscriminator::new(50, vec![25], n_tasks);
let mut task_loss_curves = HashMap::new();
for task_name in self.task_outputs.keys() {
task_loss_curves.insert(task_name.clone(), vec![0.0; self.max_iter]);
}
let trained_state = AdversarialMultiTaskNetworkTrained {
shared_weights,
shared_biases,
private_weights,
private_biases,
output_weights,
output_biases,
task_discriminator,
n_features,
task_outputs: self.task_outputs.clone(),
shared_layer_sizes: self.config.shared_layer_sizes.clone(),
private_layer_sizes: self.config.private_layer_sizes.clone(),
shared_activation: self.shared_activation,
private_activation: self.private_activation,
output_activations: self.output_activations.clone(),
task_loss_curves,
adversarial_loss_curve: vec![0.0; self.max_iter],
orthogonality_loss_curve: vec![0.0; self.max_iter],
combined_loss_curve: vec![0.0; self.max_iter],
discriminator_accuracy_curve: vec![0.0; self.max_iter],
adversarial_strategy: self.config.adversarial_strategy,
adversarial_weight: self.config.adversarial_weight,
orthogonality_weight: self.config.orthogonality_weight,
gradient_reversal_config: self.config.gradient_reversal_config.clone(),
n_iter: self.max_iter,
};
Ok(AdversarialMultiTaskNetwork {
state: trained_state,
config: self.config,
task_outputs: self.task_outputs,
task_loss_functions: self.task_loss_functions,
task_weights: self.task_weights,
shared_activation: self.shared_activation,
private_activation: self.private_activation,
output_activations: self.output_activations,
learning_rate: self.learning_rate,
max_iter: self.max_iter,
tolerance: self.tolerance,
random_state: self.random_state,
alpha: self.alpha,
})
}
}
impl Predict<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
for AdversarialMultiTaskNetwork<AdversarialMultiTaskNetworkTrained>
{
fn predict(&self, X: &ArrayView2<'_, Float>) -> SklResult<HashMap<String, Array2<Float>>> {
let (n_samples, n_features) = X.dim();
if n_features != self.state.n_features {
return Err(SklearsError::InvalidInput(
"X has different number of features than training data".to_string(),
));
}
let mut predictions = HashMap::new();
for (task_name, &output_size) in &self.state.task_outputs {
let task_pred = Array2::<Float>::zeros((n_samples, output_size));
predictions.insert(task_name.clone(), task_pred);
}
Ok(predictions)
}
}
impl AdversarialMultiTaskNetwork<AdversarialMultiTaskNetworkTrained> {
pub fn task_loss_curves(&self) -> &HashMap<String, Vec<Float>> {
&self.state.task_loss_curves
}
pub fn adversarial_loss_curve(&self) -> &[Float] {
&self.state.adversarial_loss_curve
}
pub fn orthogonality_loss_curve(&self) -> &[Float] {
&self.state.orthogonality_loss_curve
}
pub fn combined_loss_curve(&self) -> &[Float] {
&self.state.combined_loss_curve
}
pub fn discriminator_accuracy_curve(&self) -> &[Float] {
&self.state.discriminator_accuracy_curve
}
pub fn n_iter(&self) -> usize {
self.state.n_iter
}
}