Skip to main content

sklears_multioutput/
regularization.rs

1//! Multi-Task Regularization Methods
2//!
3//! This module provides various regularization techniques specifically designed for multi-task
4//! and multi-output learning scenarios. These methods help in learning shared structure
5//! across tasks while preventing overfitting.
6//!
7//! The module has been refactored into smaller submodules to comply with the 2000-line limit:
8//!
9//! - [`simd_ops`] - SIMD-accelerated operations for high-performance regularization computations
10//! - [`group_lasso`] - Group Lasso regularization for feature group selection
11//! - [`nuclear_norm`] - Nuclear norm regularization for low-rank structure learning
12//! - [`task_clustering`] - Task clustering regularization for similar task grouping
13//! - [`task_relationship`] - Task relationship learning for explicit task relationships
14//! - [`meta_learning`] - Meta-learning approach for quick adaptation to new tasks
15
16// Use SciRS2-Core for arrays and random number generation (SciRS2 Policy)
17use scirs2_core::ndarray::{Array1, Array2};
18use sklears_core::{traits::Untrained, types::Float};
19use std::collections::HashMap;
20
21// Submodules
22#[path = "regularization/simd_ops.rs"]
23pub mod simd_ops;
24
25#[path = "regularization/group_lasso.rs"]
26pub mod group_lasso;
27
28#[path = "regularization/nuclear_norm.rs"]
29pub mod nuclear_norm;
30
31#[path = "regularization/task_clustering.rs"]
32pub mod task_clustering;
33
34#[path = "regularization/task_relationship.rs"]
35pub mod task_relationship;
36
37#[path = "regularization/meta_learning.rs"]
38pub mod meta_learning;
39
40// Re-export the main types from submodules
41pub use group_lasso::{GroupLasso, GroupLassoTrained};
42pub use meta_learning::{MetaLearningMultiTask, MetaLearningMultiTaskTrained};
43pub use nuclear_norm::{NuclearNormRegression, NuclearNormRegressionTrained};
44pub use task_clustering::{TaskClusteringRegressionTrained, TaskClusteringRegularization};
45pub use task_relationship::{
46    TaskRelationshipLearning, TaskRelationshipLearningTrained, TaskSimilarityMethod,
47};
48
49/// Multi-Task Elastic Net with Group Structure
50///
51/// Combines L1 and L2 regularization with group structure awareness.
52/// Useful for scenarios where we want both feature selection and group selection.
53///
54/// Note: This is a placeholder struct - implementation is not yet complete.
55#[derive(Debug, Clone)]
56pub struct MultiTaskElasticNet<S = Untrained> {
57    #[allow(dead_code)]
58    state: S,
59    #[allow(dead_code)]
60    /// L1 regularization strength
61    alpha: Float,
62    #[allow(dead_code)]
63    /// L1 vs L2 balance (0 = Ridge, 1 = Lasso)
64    l1_ratio: Float,
65    #[allow(dead_code)]
66    /// Feature groups for group penalties
67    feature_groups: Vec<Vec<usize>>,
68    #[allow(dead_code)]
69    /// Group penalty strength
70    group_alpha: Float,
71    #[allow(dead_code)]
72    /// Maximum number of iterations
73    max_iter: usize,
74    #[allow(dead_code)]
75    /// Convergence tolerance
76    tolerance: Float,
77    #[allow(dead_code)]
78    /// Learning rate
79    learning_rate: Float,
80    #[allow(dead_code)]
81    /// Task configurations
82    task_outputs: HashMap<String, usize>,
83    #[allow(dead_code)]
84    /// Include intercept term
85    fit_intercept: bool,
86}
87
88/// Trained state for MultiTaskElasticNet
89///
90/// Note: This is a placeholder struct - implementation is not yet complete.
91#[derive(Debug, Clone)]
92pub struct MultiTaskElasticNetTrained {
93    #[allow(dead_code)]
94    /// Coefficients for each task
95    coefficients: HashMap<String, Array2<Float>>,
96    #[allow(dead_code)]
97    /// Intercepts for each task
98    intercepts: HashMap<String, Array1<Float>>,
99    #[allow(dead_code)]
100    /// Number of input features
101    n_features: usize,
102    #[allow(dead_code)]
103    /// Task configurations
104    task_outputs: HashMap<String, usize>,
105    #[allow(dead_code)]
106    /// Training parameters
107    alpha: Float,
108    #[allow(dead_code)]
109    l1_ratio: Float,
110    #[allow(dead_code)]
111    group_alpha: Float,
112    #[allow(dead_code)]
113    /// Training iterations performed
114    n_iter: usize,
115}
116
117/// Regularization strategies for multi-task learning
118#[derive(Debug, Clone, PartialEq, Default)]
119pub enum RegularizationStrategy {
120    /// No regularization
121    #[default]
122    None,
123    /// L1 regularization (Lasso)
124    L1(Float),
125    /// L2 regularization (Ridge)
126    L2(Float),
127    /// Elastic Net (L1 + L2)
128    ElasticNet { alpha: Float, l1_ratio: Float },
129    /// Group Lasso
130    GroupLasso { alpha: Float },
131    /// Nuclear norm regularization
132    NuclearNorm { alpha: Float },
133    /// Task clustering regularization
134    TaskClustering {
135        n_clusters: usize,
136        intra_cluster_alpha: Float,
137        inter_cluster_alpha: Float,
138    },
139    /// Task relationship learning
140    TaskRelationship {
141        relationship_strength: Float,
142        similarity_threshold: Float,
143    },
144    /// Meta-learning for multi-task
145    MetaLearning {
146        meta_learning_rate: Float,
147        inner_learning_rate: Float,
148        n_inner_steps: usize,
149    },
150}
151
152// Keep the tests in the main module for backwards compatibility
153#[allow(non_snake_case)]
154#[cfg(test)]
155mod regularization_tests {
156    use super::*;
157    use approx::assert_abs_diff_eq;
158    // Use SciRS2-Core for arrays and random number generation (SciRS2 Policy)
159    use scirs2_core::ndarray::array;
160    use sklears_core::traits::{Fit, Predict};
161    use std::collections::HashMap;
162
163    #[test]
164    fn test_group_lasso_creation() {
165        let group_lasso = GroupLasso::new()
166            .alpha(0.1)
167            .feature_groups(vec![vec![0, 1], vec![2, 3]])
168            .max_iter(100)
169            .tolerance(1e-6)
170            .learning_rate(0.01);
171
172        assert_eq!(group_lasso.alpha, 0.1);
173        assert_eq!(group_lasso.feature_groups, vec![vec![0, 1], vec![2, 3]]);
174        assert_eq!(group_lasso.max_iter, 100);
175        assert_abs_diff_eq!(group_lasso.tolerance, 1e-6);
176        assert_abs_diff_eq!(group_lasso.learning_rate, 0.01);
177    }
178
179    #[test]
180    fn test_group_lasso_fit_predict() {
181        let X = array![
182            [1.0, 2.0, 3.0, 4.0],
183            [2.0, 3.0, 4.0, 5.0],
184            [3.0, 1.0, 2.0, 3.0],
185            [4.0, 2.0, 1.0, 2.0]
186        ];
187
188        let mut y_tasks = HashMap::new();
189        y_tasks.insert("task1".to_string(), array![[1.0], [2.0], [1.5], [2.5]]);
190        y_tasks.insert("task2".to_string(), array![[0.5], [1.0], [0.8], [1.2]]);
191
192        let feature_groups = vec![vec![0, 1], vec![2, 3]];
193
194        let group_lasso = GroupLasso::new()
195            .alpha(0.01)
196            .feature_groups(feature_groups)
197            .task_outputs(&[("task1", 1), ("task2", 1)])
198            .max_iter(50)
199            .tolerance(1e-4)
200            .learning_rate(0.01);
201
202        let trained = group_lasso
203            .fit(&X.view(), &y_tasks)
204            .expect("model fitting should succeed");
205
206        // Test predictions
207        let predictions = trained
208            .predict(&X.view())
209            .expect("prediction should succeed");
210        assert!(predictions.contains_key("task1"));
211        assert!(predictions.contains_key("task2"));
212
213        let task1_pred = &predictions["task1"];
214        let task2_pred = &predictions["task2"];
215
216        assert_eq!(task1_pred.shape(), &[4, 1]);
217        assert_eq!(task2_pred.shape(), &[4, 1]);
218
219        // Test group sparsity
220        let sparsity = trained.group_sparsity();
221        assert!((0.0..=1.0).contains(&sparsity)); // Should be a percentage
222
223        // Test accessors
224        assert!(trained.task_coefficients("task1").is_some());
225        assert!(trained.task_intercepts("task1").is_some());
226        assert!(trained.n_iter() <= 50);
227    }
228
229    #[test]
230    fn test_nuclear_norm_regression_creation() {
231        let nuclear_norm = NuclearNormRegression::new()
232            .alpha(0.1)
233            .max_iter(100)
234            .tolerance(1e-6)
235            .learning_rate(0.01)
236            .target_rank(Some(5));
237
238        assert_eq!(nuclear_norm.alpha, 0.1);
239        assert_eq!(nuclear_norm.max_iter, 100);
240        assert_abs_diff_eq!(nuclear_norm.tolerance, 1e-6);
241        assert_abs_diff_eq!(nuclear_norm.learning_rate, 0.01);
242        assert_eq!(nuclear_norm.target_rank, Some(5));
243    }
244
245    #[test]
246    fn test_nuclear_norm_regression_fit_predict() {
247        let X = array![[1.0, 2.0], [2.0, 3.0], [3.0, 1.0], [4.0, 4.0]];
248
249        let mut y_tasks = HashMap::new();
250        y_tasks.insert("task1".to_string(), array![[1.0], [2.0], [1.5], [2.5]]);
251        y_tasks.insert("task2".to_string(), array![[0.5], [1.0], [0.8], [1.2]]);
252
253        let nuclear_norm = NuclearNormRegression::new()
254            .alpha(0.01)
255            .task_outputs(&[("task1", 1), ("task2", 1)])
256            .max_iter(50)
257            .tolerance(1e-4)
258            .learning_rate(0.01);
259
260        let trained = nuclear_norm
261            .fit(&X.view(), &y_tasks)
262            .expect("model fitting should succeed");
263
264        // Test predictions
265        let predictions = trained
266            .predict(&X.view())
267            .expect("prediction should succeed");
268        assert!(predictions.contains_key("task1"));
269        assert!(predictions.contains_key("task2"));
270
271        let task1_pred = &predictions["task1"];
272        let task2_pred = &predictions["task2"];
273
274        assert_eq!(task1_pred.shape(), &[4, 1]);
275        assert_eq!(task2_pred.shape(), &[4, 1]);
276
277        // Test accessors
278        assert!(trained.task_coefficient_matrix("task1").is_some());
279        assert!(trained.effective_rank() < usize::MAX); // effective_rank is always non-negative (usize)
280        assert!(!trained.singular_values().is_empty());
281        assert!(trained.n_iter() <= 50);
282    }
283
284    #[test]
285    fn test_regularization_strategies() {
286        let strategies = vec![
287            RegularizationStrategy::None,
288            RegularizationStrategy::L1(0.1),
289            RegularizationStrategy::L2(0.1),
290            RegularizationStrategy::ElasticNet {
291                alpha: 0.1,
292                l1_ratio: 0.5,
293            },
294            RegularizationStrategy::GroupLasso { alpha: 0.1 },
295            RegularizationStrategy::NuclearNorm { alpha: 0.1 },
296            RegularizationStrategy::TaskClustering {
297                n_clusters: 2,
298                intra_cluster_alpha: 0.1,
299                inter_cluster_alpha: 0.01,
300            },
301            RegularizationStrategy::TaskRelationship {
302                relationship_strength: 0.1,
303                similarity_threshold: 0.5,
304            },
305            RegularizationStrategy::MetaLearning {
306                meta_learning_rate: 0.01,
307                inner_learning_rate: 0.1,
308                n_inner_steps: 5,
309            },
310        ];
311
312        assert_eq!(strategies.len(), 9);
313        assert_eq!(strategies[0], RegularizationStrategy::None);
314        assert_eq!(strategies[1], RegularizationStrategy::L1(0.1));
315    }
316
317    #[test]
318    fn test_task_similarity_methods() {
319        let methods = [
320            TaskSimilarityMethod::Correlation,
321            TaskSimilarityMethod::Cosine,
322            TaskSimilarityMethod::Euclidean,
323            TaskSimilarityMethod::MutualInformation,
324        ];
325
326        assert_eq!(methods.len(), 4);
327        assert_eq!(methods[0], TaskSimilarityMethod::Correlation);
328        assert_eq!(methods[1], TaskSimilarityMethod::Cosine);
329    }
330}