1use scirs2_core::ndarray::{Array1, Array2};
18use sklears_core::{traits::Untrained, types::Float};
19use std::collections::HashMap;
20
21#[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
40pub 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#[derive(Debug, Clone)]
56pub struct MultiTaskElasticNet<S = Untrained> {
57 #[allow(dead_code)]
58 state: S,
59 #[allow(dead_code)]
60 alpha: Float,
62 #[allow(dead_code)]
63 l1_ratio: Float,
65 #[allow(dead_code)]
66 feature_groups: Vec<Vec<usize>>,
68 #[allow(dead_code)]
69 group_alpha: Float,
71 #[allow(dead_code)]
72 max_iter: usize,
74 #[allow(dead_code)]
75 tolerance: Float,
77 #[allow(dead_code)]
78 learning_rate: Float,
80 #[allow(dead_code)]
81 task_outputs: HashMap<String, usize>,
83 #[allow(dead_code)]
84 fit_intercept: bool,
86}
87
88#[derive(Debug, Clone)]
92pub struct MultiTaskElasticNetTrained {
93 #[allow(dead_code)]
94 coefficients: HashMap<String, Array2<Float>>,
96 #[allow(dead_code)]
97 intercepts: HashMap<String, Array1<Float>>,
99 #[allow(dead_code)]
100 n_features: usize,
102 #[allow(dead_code)]
103 task_outputs: HashMap<String, usize>,
105 #[allow(dead_code)]
106 alpha: Float,
108 #[allow(dead_code)]
109 l1_ratio: Float,
110 #[allow(dead_code)]
111 group_alpha: Float,
112 #[allow(dead_code)]
113 n_iter: usize,
115}
116
117#[derive(Debug, Clone, PartialEq, Default)]
119pub enum RegularizationStrategy {
120 #[default]
122 None,
123 L1(Float),
125 L2(Float),
127 ElasticNet { alpha: Float, l1_ratio: Float },
129 GroupLasso { alpha: Float },
131 NuclearNorm { alpha: Float },
133 TaskClustering {
135 n_clusters: usize,
136 intra_cluster_alpha: Float,
137 inter_cluster_alpha: Float,
138 },
139 TaskRelationship {
141 relationship_strength: Float,
142 similarity_threshold: Float,
143 },
144 MetaLearning {
146 meta_learning_rate: Float,
147 inner_learning_rate: Float,
148 n_inner_steps: usize,
149 },
150}
151
152#[allow(non_snake_case)]
154#[cfg(test)]
155mod regularization_tests {
156 use super::*;
157 use approx::assert_abs_diff_eq;
158 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 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 let sparsity = trained.group_sparsity();
221 assert!((0.0..=1.0).contains(&sparsity)); 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 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 assert!(trained.task_coefficient_matrix("task1").is_some());
279 assert!(trained.effective_rank() < usize::MAX); 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}