pub use crate::activation::ActivationFunction;
pub use crate::loss::LossFunction;
pub use crate::mlp::{
MultiOutputMLP, MultiOutputMLPClassifier, MultiOutputMLPRegressor, MultiOutputMLPTrained,
};
pub use crate::recurrent::{
CellType, RecurrentNeuralNetwork, RecurrentNeuralNetworkTrained, SequenceMode,
};
pub use crate::multitask::{MultiTaskNeuralNetwork, MultiTaskNeuralNetworkTrained, TaskBalancing};
pub use crate::adversarial::{
AdversarialMultiTaskNetwork, AdversarialMultiTaskNetworkTrained, AdversarialStrategy,
GradientReversalConfig, LambdaSchedule, TaskDiscriminator,
};
#[allow(non_snake_case)]
#[cfg(test)]
mod tests {
use super::*;
use scirs2_core::ndarray::array;
use sklears_core::traits::{Fit, Predict};
use std::collections::HashMap;
#[test]
fn test_activation_functions_integration() {
let x = array![1.0, -1.0, 0.0, 2.0];
let relu_result = ActivationFunction::ReLU.apply(&x);
assert_eq!(relu_result, array![1.0, 0.0, 0.0, 2.0]);
let linear_result = ActivationFunction::Linear.apply(&x);
assert_eq!(linear_result, x);
}
#[test]
fn test_mlp_integration() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let y = array![[0.5, 1.2], [1.0, 2.1]];
let mlp = MultiOutputMLP::new()
.hidden_layer_sizes(vec![5])
.max_iter(5)
.random_state(Some(42));
let trained_mlp = mlp
.fit(&X.view(), &y)
.expect("model fitting should succeed");
let predictions = trained_mlp
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.dim(), (2, 2));
assert!(!predictions.iter().any(|&x| x.is_nan()));
}
#[test]
fn test_multitask_integration() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let mut tasks = HashMap::new();
tasks.insert("task1".to_string(), array![[0.5], [1.0]]);
tasks.insert("task2".to_string(), array![[1.0], [0.0]]);
let mt_net = MultiTaskNeuralNetwork::new()
.task_outputs(&[("task1", 1), ("task2", 1)])
.max_iter(5)
.random_state(Some(42));
let trained = mt_net
.fit(&X.view(), &tasks)
.expect("model fitting should succeed");
let predictions = trained
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.len(), 2);
assert!(predictions.contains_key("task1"));
assert!(predictions.contains_key("task2"));
}
#[test]
fn test_adversarial_integration() {
let X = array![[1.0, 2.0], [2.0, 3.0]];
let mut tasks = HashMap::new();
tasks.insert("task1".to_string(), array![[0.5], [1.0]]);
let adv_net = AdversarialMultiTaskNetwork::new()
.task_outputs(&[("task1", 1)])
.max_iter(5)
.random_state(Some(42));
let trained = adv_net
.fit(&X.view(), &tasks)
.expect("model fitting should succeed");
let predictions = trained
.predict(&X.view())
.expect("prediction should succeed");
assert_eq!(predictions.len(), 1);
assert!(predictions.contains_key("task1"));
}
#[test]
fn test_module_exports() {
use super::{
ActivationFunction, AdversarialMultiTaskNetwork, AdversarialStrategy, CellType,
LossFunction, MultiOutputMLP, MultiTaskNeuralNetwork, RecurrentNeuralNetwork,
SequenceMode, TaskBalancing,
};
let _mlp = MultiOutputMLP::new();
let _rnn = RecurrentNeuralNetwork::new();
let _mt = MultiTaskNeuralNetwork::new();
let _adv = AdversarialMultiTaskNetwork::new();
let _act = ActivationFunction::ReLU;
let _loss = LossFunction::MeanSquaredError;
let _cell = CellType::LSTM;
let _seq = SequenceMode::ManyToMany;
let _bal = TaskBalancing::Equal;
let _strat = AdversarialStrategy::GradientReversal;
}
}