optirs_core/regularizers/
manifold.rs1use scirs2_core::ndarray::{Array, Array2, ArrayBase, Data, Dimension, ScalarOperand};
7use scirs2_core::numeric::{Float, FromPrimitive};
8use std::fmt::Debug;
9
10use crate::error::{OptimError, Result};
11use crate::regularizers::Regularizer;
12
13#[derive(Debug, Clone)]
37pub struct ManifoldRegularization<A: Float> {
38 lambda: A,
40 similarity_matrix: Option<Array2<A>>,
42 degree_matrix: Option<Array2<A>>,
44 laplacian: Option<Array2<A>>,
46}
47
48impl<A: Float + Debug + ScalarOperand + FromPrimitive + Send + Sync> ManifoldRegularization<A> {
49 pub fn new(lambda: A) -> Self {
55 Self {
56 lambda,
57 similarity_matrix: None,
58 degree_matrix: None,
59 laplacian: None,
60 }
61 }
62
63 pub fn set_similarity_matrix(&mut self, similarity: Array2<A>) -> Result<()> {
69 let (rows, cols) = similarity.dim();
70 if rows != cols {
71 return Err(OptimError::InvalidConfig(
72 "Similarity matrix must be square".to_string(),
73 ));
74 }
75
76 let mut degree = Array2::zeros((rows, rows));
78 for i in 0..rows {
79 let row_sum = similarity.row(i).sum();
80 degree[[i, i]] = row_sum;
81 }
82
83 let laplacian = °ree - &similarity;
85
86 self.similarity_matrix = Some(similarity);
87 self.degree_matrix = Some(degree);
88 self.laplacian = Some(laplacian);
89
90 Ok(())
91 }
92
93 pub fn compute_penalty<S>(&self, params: &ArrayBase<S, scirs2_core::ndarray::Ix2>) -> Result<A>
95 where
96 S: Data<Elem = A>,
97 {
98 let laplacian = self
99 .laplacian
100 .as_ref()
101 .ok_or_else(|| OptimError::InvalidConfig("Similarity matrix not set".to_string()))?;
102
103 let lf = laplacian.dot(params);
106 let penalty = params
107 .iter()
108 .zip(lf.iter())
109 .map(|(p, lf)| *p * *lf)
110 .fold(A::zero(), |acc, val| acc + val);
111
112 Ok(self.lambda * penalty)
113 }
114
115 fn compute_gradient<S>(
117 &self,
118 params: &ArrayBase<S, scirs2_core::ndarray::Ix2>,
119 ) -> Result<Array2<A>>
120 where
121 S: Data<Elem = A>,
122 {
123 let laplacian = self
124 .laplacian
125 .as_ref()
126 .ok_or_else(|| OptimError::InvalidConfig("Similarity matrix not set".to_string()))?;
127
128 let two: A = crate::regularizers::cast_scalar(2.0)?;
130 let gradient = laplacian.dot(params) * (two * self.lambda);
131 Ok(gradient)
132 }
133}
134
135impl<
137 A: Float + Debug + ScalarOperand + FromPrimitive + Send + Sync,
138 D: Dimension + Send + Sync,
139 > Regularizer<A, D> for ManifoldRegularization<A>
140{
141 fn apply(&self, params: &Array<A, D>, gradients: &mut Array<A, D>) -> Result<A> {
142 if params.ndim() != 2 {
143 return Ok(A::zero());
145 }
146
147 let params_2d = params
149 .view()
150 .into_dimensionality::<scirs2_core::ndarray::Ix2>()
151 .map_err(|_| OptimError::InvalidConfig("Expected 2D array".to_string()))?;
152
153 let gradient_update = self.compute_gradient(¶ms_2d)?;
154
155 let mut gradients_2d = gradients
157 .view_mut()
158 .into_dimensionality::<scirs2_core::ndarray::Ix2>()
159 .map_err(|_| OptimError::InvalidConfig("Expected 2D array".to_string()))?;
160
161 gradients_2d.zip_mut_with(&gradient_update, |g, &u| *g = *g + u);
162
163 self.compute_penalty(¶ms_2d)
165 }
166
167 fn penalty(&self, params: &Array<A, D>) -> Result<A> {
168 if params.ndim() != 2 {
169 return Ok(A::zero());
171 }
172
173 let params_2d = params
175 .view()
176 .into_dimensionality::<scirs2_core::ndarray::Ix2>()
177 .map_err(|_| OptimError::InvalidConfig("Expected 2D array".to_string()))?;
178
179 self.compute_penalty(¶ms_2d)
180 }
181}
182
183#[cfg(test)]
184mod tests {
185 use super::*;
186 use approx::assert_relative_eq;
187 use scirs2_core::ndarray::array;
188
189 #[test]
190 fn test_manifold_creation() {
191 let manifold = ManifoldRegularization::<f64>::new(0.01);
192 assert_eq!(manifold.lambda, 0.01);
193 assert!(manifold.similarity_matrix.is_none());
194 }
195
196 #[test]
197 fn test_set_similarity_matrix() {
198 let mut manifold = ManifoldRegularization::new(0.01);
199
200 let similarity = array![[1.0, 0.5], [0.5, 1.0]];
202
203 assert!(manifold.set_similarity_matrix(similarity).is_ok());
204 assert!(manifold.laplacian.is_some());
205
206 let laplacian = manifold
208 .laplacian
209 .as_ref()
210 .expect("manifold.laplacian.as_ref succeeds in test_set_similarity_matrix");
211 assert_relative_eq!(laplacian[[0, 0]], 0.5, epsilon = 1e-10);
215 assert_relative_eq!(laplacian[[0, 1]], -0.5, epsilon = 1e-10);
216 }
217
218 #[test]
219 fn test_invalid_similarity_matrix() {
220 let mut manifold = ManifoldRegularization::<f64>::new(0.01);
221
222 let similarity = array![[1.0, 0.5, 0.3], [0.5, 1.0, 0.4]];
224 assert!(manifold.set_similarity_matrix(similarity).is_err());
225 }
226
227 #[test]
228 fn test_penalty_without_similarity() {
229 let manifold = ManifoldRegularization::<f64>::new(0.01);
230 let params = array![[1.0, 2.0], [3.0, 4.0]];
231
232 assert!(manifold.compute_penalty(¶ms).is_err());
234 }
235
236 #[test]
237 fn test_penalty_computation() {
238 let mut manifold = ManifoldRegularization::new(0.1);
239
240 let similarity = array![[1.0, 0.8], [0.8, 1.0]];
242 manifold
243 .set_similarity_matrix(similarity)
244 .expect("set_similarity_matrix succeeds in test_penalty_computation");
245
246 let params = array![[1.0, 0.0], [0.0, 1.0]];
248 let penalty = manifold
249 .compute_penalty(¶ms)
250 .expect("manifold.compute_penalty succeeds in test_penalty_computation");
251
252 assert!(penalty > 0.0);
254 }
255
256 #[test]
257 fn test_gradient_computation() {
258 let mut manifold = ManifoldRegularization::new(0.1);
259
260 let similarity = array![[1.0, 0.8], [0.8, 1.0]];
262 manifold
263 .set_similarity_matrix(similarity)
264 .expect("set_similarity_matrix succeeds in test_gradient_computation");
265
266 let params = array![[1.0, 2.0], [3.0, 4.0]];
267 let gradient = manifold
268 .compute_gradient(¶ms)
269 .expect("manifold.compute_gradient succeeds in test_gradient_computation");
270
271 assert!(gradient.abs().sum() > 0.0);
273 }
274
275 #[test]
276 fn test_regularizer_trait() {
277 let mut manifold = ManifoldRegularization::new(0.01);
278
279 let similarity = array![[1.0, 0.6], [0.6, 1.0]];
281 manifold
282 .set_similarity_matrix(similarity)
283 .expect("set_similarity_matrix succeeds in test_regularizer_trait");
284
285 let params = array![[1.0, 2.0], [3.0, 4.0]];
286 let mut gradient = array![[0.1, 0.2], [0.3, 0.4]];
287 let original_gradient = gradient.clone();
288
289 let penalty = manifold
290 .apply(¶ms, &mut gradient)
291 .expect("apply succeeds in test_regularizer_trait");
292
293 assert!(penalty > 0.0);
295
296 assert_ne!(gradient, original_gradient);
298 }
299
300 #[test]
301 fn test_identity_similarity() {
302 let mut manifold = ManifoldRegularization::new(0.1);
303
304 let similarity = array![[1.0, 0.0], [0.0, 1.0]];
306 manifold
307 .set_similarity_matrix(similarity)
308 .expect("set_similarity_matrix succeeds in test_identity_similarity");
309
310 let params = array![[1.0, 2.0], [3.0, 4.0]];
311 let penalty = manifold
312 .compute_penalty(¶ms)
313 .expect("manifold.compute_penalty succeeds in test_identity_similarity");
314
315 assert!(penalty >= 0.0);
317 }
318}