#![allow(non_snake_case)]
use scirs2_core::ndarray::{Array1, Array2, ArrayView2};
use scirs2_core::random::thread_rng;
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, PartialEq)]
pub enum TaskBalancing {
Equal,
Weighted,
Adaptive,
GradientBalancing,
}
#[derive(Debug, Clone)]
pub struct MultiTaskNeuralNetwork<S = Untrained> {
state: S,
shared_layer_sizes: Vec<usize>,
task_specific_layer_sizes: Vec<usize>,
task_outputs: HashMap<String, usize>,
task_loss_functions: HashMap<String, LossFunction>,
task_weights: HashMap<String, Float>,
shared_activation: ActivationFunction,
task_activation: ActivationFunction,
output_activations: HashMap<String, ActivationFunction>,
learning_rate: Float,
max_iter: usize,
tolerance: Float,
random_state: Option<u64>,
alpha: Float,
batch_size: Option<usize>,
early_stopping: bool,
validation_fraction: Float,
task_balancing: TaskBalancing,
}
#[derive(Debug, Clone)]
pub struct MultiTaskNeuralNetworkTrained {
#[allow(dead_code)]
shared_weights: Vec<Array2<Float>>,
#[allow(dead_code)]
shared_biases: Vec<Array1<Float>>,
#[allow(dead_code)]
task_weights: HashMap<String, Vec<Array2<Float>>>,
#[allow(dead_code)]
task_biases: HashMap<String, Vec<Array1<Float>>>,
#[allow(dead_code)]
output_weights: HashMap<String, Array2<Float>>,
#[allow(dead_code)]
output_biases: HashMap<String, Array1<Float>>,
n_features: usize,
task_outputs: HashMap<String, usize>,
#[allow(dead_code)]
shared_layer_sizes: Vec<usize>,
#[allow(dead_code)]
task_specific_layer_sizes: Vec<usize>,
#[allow(dead_code)]
shared_activation: ActivationFunction,
#[allow(dead_code)]
task_activation: ActivationFunction,
#[allow(dead_code)]
output_activations: HashMap<String, ActivationFunction>,
task_loss_curves: HashMap<String, Vec<Float>>,
combined_loss_curve: Vec<Float>,
n_iter: usize,
}
impl MultiTaskNeuralNetwork<Untrained> {
pub fn new() -> Self {
Self {
state: Untrained,
shared_layer_sizes: vec![100],
task_specific_layer_sizes: vec![50],
task_outputs: HashMap::new(),
task_loss_functions: HashMap::new(),
task_weights: HashMap::new(),
shared_activation: ActivationFunction::ReLU,
task_activation: ActivationFunction::ReLU,
output_activations: HashMap::new(),
learning_rate: 0.001,
max_iter: 1000,
tolerance: 1e-6,
random_state: None,
alpha: 0.0001,
batch_size: None,
early_stopping: false,
validation_fraction: 0.1,
task_balancing: TaskBalancing::Equal,
}
}
pub fn shared_layers(mut self, sizes: Vec<usize>) -> Self {
self.shared_layer_sizes = sizes;
self
}
pub fn task_specific_layers(mut self, sizes: Vec<usize>) -> Self {
self.task_specific_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 task_loss_functions(mut self, loss_functions: &[(&str, LossFunction)]) -> Self {
for (task_name, loss_fn) in loss_functions {
self.task_loss_functions
.insert(task_name.to_string(), *loss_fn);
}
self
}
pub fn task_weights(mut self, weights: &[(&str, Float)]) -> Self {
for (task_name, weight) in weights {
self.task_weights.insert(task_name.to_string(), *weight);
}
self
}
pub fn shared_activation(mut self, activation: ActivationFunction) -> Self {
self.shared_activation = activation;
self
}
pub fn task_activation(mut self, activation: ActivationFunction) -> Self {
self.task_activation = activation;
self
}
pub fn output_activations(mut self, activations: &[(&str, ActivationFunction)]) -> Self {
for (task_name, activation) in activations {
self.output_activations
.insert(task_name.to_string(), *activation);
}
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 tolerance(mut self, tolerance: Float) -> Self {
self.tolerance = tolerance;
self
}
pub fn random_state(mut self, seed: Option<u64>) -> Self {
self.random_state = seed;
self
}
pub fn alpha(mut self, alpha: Float) -> Self {
self.alpha = alpha;
self
}
pub fn batch_size(mut self, batch_size: Option<usize>) -> Self {
self.batch_size = batch_size;
self
}
pub fn early_stopping(mut self, early_stopping: bool) -> Self {
self.early_stopping = early_stopping;
self
}
pub fn validation_fraction(mut self, fraction: Float) -> Self {
self.validation_fraction = fraction;
self
}
pub fn task_balancing(mut self, strategy: TaskBalancing) -> Self {
self.task_balancing = strategy;
self
}
}
impl Default for MultiTaskNeuralNetwork<Untrained> {
fn default() -> Self {
Self::new()
}
}
impl Estimator for MultiTaskNeuralNetwork<Untrained> {
type Config = ();
type Error = SklearsError;
type Float = Float;
fn config(&self) -> &Self::Config {
&()
}
}
impl Fit<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
for MultiTaskNeuralNetwork<Untrained>
{
type Fitted = MultiTaskNeuralNetwork<MultiTaskNeuralNetworkTrained>;
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_samples = x.nrows();
for (task_name, task_targets) in y {
if task_targets.nrows() != n_samples {
return Err(SklearsError::ShapeMismatch {
expected: format!("{}", n_samples),
actual: format!("{}", task_targets.nrows()),
});
}
if !self.task_outputs.contains_key(task_name) {
return Err(SklearsError::InvalidInput(format!(
"Unknown task: {}",
task_name
)));
}
}
let n_features = x.ncols();
let _rng = thread_rng();
let shared_weights = vec![Array2::<Float>::zeros((n_features, 50))];
let shared_biases = vec![Array1::<Float>::zeros(50)];
let mut task_weights = HashMap::new();
let mut task_biases = HashMap::new();
let mut output_weights = HashMap::new();
let mut output_biases = HashMap::new();
for (task_name, &output_size) in &self.task_outputs {
task_weights.insert(task_name.clone(), vec![Array2::<Float>::zeros((50, 25))]);
task_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 mut task_loss_curves = HashMap::new();
let combined_loss_curve = vec![0.0; self.max_iter];
for task_name in self.task_outputs.keys() {
task_loss_curves.insert(task_name.clone(), vec![0.0; self.max_iter]);
}
let trained_state = MultiTaskNeuralNetworkTrained {
shared_weights,
shared_biases,
task_weights,
task_biases,
output_weights,
output_biases,
n_features,
task_outputs: self.task_outputs.clone(),
shared_layer_sizes: self.shared_layer_sizes.clone(),
task_specific_layer_sizes: self.task_specific_layer_sizes.clone(),
shared_activation: self.shared_activation,
task_activation: self.task_activation,
output_activations: self.output_activations.clone(),
task_loss_curves,
combined_loss_curve,
n_iter: self.max_iter,
};
Ok(MultiTaskNeuralNetwork {
state: trained_state,
shared_layer_sizes: self.shared_layer_sizes,
task_specific_layer_sizes: self.task_specific_layer_sizes,
task_outputs: self.task_outputs,
task_loss_functions: self.task_loss_functions,
task_weights: self.task_weights,
shared_activation: self.shared_activation,
task_activation: self.task_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,
batch_size: self.batch_size,
early_stopping: self.early_stopping,
validation_fraction: self.validation_fraction,
task_balancing: self.task_balancing,
})
}
}
impl Predict<ArrayView2<'_, Float>, HashMap<String, Array2<Float>>>
for MultiTaskNeuralNetwork<MultiTaskNeuralNetworkTrained>
{
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 MultiTaskNeuralNetwork<MultiTaskNeuralNetworkTrained> {
pub fn task_loss_curves(&self) -> &HashMap<String, Vec<Float>> {
&self.state.task_loss_curves
}
pub fn combined_loss_curve(&self) -> &[Float] {
&self.state.combined_loss_curve
}
pub fn n_iter(&self) -> usize {
self.state.n_iter
}
pub fn task_outputs(&self) -> &HashMap<String, usize> {
&self.state.task_outputs
}
}