optirs_core/regularizers/
activity.rs1use scirs2_core::ndarray::{Array, ArrayBase, Data, Dimension, ScalarOperand};
2use scirs2_core::numeric::{Float, FromPrimitive};
3use std::fmt::Debug;
4
5use crate::error::Result;
6use crate::regularizers::Regularizer;
7
8#[derive(Debug, Clone, Copy, PartialEq)]
10pub enum ActivityNorm {
11 L1,
13 L2,
15 L2Squared,
17}
18
19#[derive(Debug, Clone, Copy)]
36pub struct ActivityRegularization<A: Float + FromPrimitive + Debug> {
37 pub lambda: A,
39 pub norm: ActivityNorm,
41}
42
43impl<A: Float + FromPrimitive + Debug + Send + Sync> ActivityRegularization<A> {
44 pub fn l1(lambda: A) -> Self {
54 Self {
55 lambda,
56 norm: ActivityNorm::L1,
57 }
58 }
59
60 pub fn l2(lambda: A) -> Self {
70 Self {
71 lambda,
72 norm: ActivityNorm::L2,
73 }
74 }
75
76 pub fn l2_squared(lambda: A) -> Self {
86 Self {
87 lambda,
88 norm: ActivityNorm::L2Squared,
89 }
90 }
91
92 pub fn new(lambda: A, norm: ActivityNorm) -> Self {
103 Self { lambda, norm }
104 }
105
106 fn calculate_penalty<S, D>(&self, activations: &ArrayBase<S, D>) -> A
116 where
117 S: Data<Elem = A>,
118 D: Dimension,
119 {
120 match self.norm {
121 ActivityNorm::L1 => {
122 let sum_abs = activations.mapv(|x| x.abs()).sum();
124 self.lambda * sum_abs
125 }
126 ActivityNorm::L2 => {
127 let sum_squared = activations.mapv(|x| x * x).sum();
129 self.lambda * sum_squared.sqrt()
130 }
131 ActivityNorm::L2Squared => {
132 let sum_squared = activations.mapv(|x| x * x).sum();
134 self.lambda * sum_squared
135 }
136 }
137 }
138
139 fn calculate_gradients<S, D>(&self, activations: &ArrayBase<S, D>) -> Array<A, D>
149 where
150 S: Data<Elem = A>,
151 D: Dimension,
152 {
153 match self.norm {
154 ActivityNorm::L1 => {
155 activations.mapv(|x| {
157 if x > A::zero() {
158 self.lambda
159 } else if x < A::zero() {
160 -self.lambda
161 } else {
162 A::zero()
163 }
164 })
165 }
166 ActivityNorm::L2 => {
167 let sum_squared = activations.mapv(|x| x * x).sum();
169
170 if sum_squared <= A::epsilon() {
172 return Array::zeros(activations.raw_dim());
173 }
174
175 let norm = sum_squared.sqrt();
176 activations.mapv(|x| self.lambda * x / norm)
177 }
178 ActivityNorm::L2Squared => {
179 let two = A::one() + A::one();
181 activations.mapv(|x| self.lambda * two * x)
182 }
183 }
184 }
185}
186
187impl<A, D> Regularizer<A, D> for ActivityRegularization<A>
188where
189 A: Float + ScalarOperand + Debug + FromPrimitive + Send + Sync,
190 D: Dimension,
191{
192 fn apply(&self, params: &Array<A, D>, gradients: &mut Array<A, D>) -> Result<A> {
193 let penalty = self.calculate_penalty(params);
195
196 let activity_grads = self.calculate_gradients(params);
198 gradients.zip_mut_with(&activity_grads, |g, &a| *g = *g + a);
199
200 Ok(penalty)
201 }
202
203 fn penalty(&self, params: &Array<A, D>) -> Result<A> {
204 Ok(self.calculate_penalty(params))
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211 use approx::assert_abs_diff_eq;
212 use scirs2_core::ndarray::array;
213 use scirs2_core::ndarray::{Array1, Array2};
214
215 #[test]
216 fn test_activity_regularization_creation() {
217 let ar = ActivityRegularization::l1(0.1f64);
218 assert_eq!(ar.lambda, 0.1);
219 assert_eq!(ar.norm, ActivityNorm::L1);
220
221 let ar = ActivityRegularization::l2(0.2f64);
222 assert_eq!(ar.lambda, 0.2);
223 assert_eq!(ar.norm, ActivityNorm::L2);
224
225 let ar = ActivityRegularization::l2_squared(0.3f64);
226 assert_eq!(ar.lambda, 0.3);
227 assert_eq!(ar.norm, ActivityNorm::L2Squared);
228
229 let ar = ActivityRegularization::new(0.4f64, ActivityNorm::L1);
230 assert_eq!(ar.lambda, 0.4);
231 assert_eq!(ar.norm, ActivityNorm::L1);
232 }
233
234 #[test]
235 fn test_l1_penalty() {
236 let lambda = 0.1f64;
237 let ar = ActivityRegularization::l1(lambda);
238
239 let activations = Array1::from_vec(vec![1.0f64, -2.0, 3.0]);
240 let penalty = ar
241 .penalty(&activations)
242 .expect("ar.penalty succeeds in test_l1_penalty");
243
244 assert_abs_diff_eq!(penalty, lambda * 6.0, epsilon = 1e-10);
246 }
247
248 #[test]
249 fn test_l2_penalty() {
250 let lambda = 0.1f64;
251 let ar = ActivityRegularization::l2(lambda);
252
253 let activations = Array1::from_vec(vec![3.0f64, 4.0]);
254 let penalty = ar
255 .penalty(&activations)
256 .expect("ar.penalty succeeds in test_l2_penalty");
257
258 assert_abs_diff_eq!(penalty, lambda * 5.0, epsilon = 1e-10);
260 }
261
262 #[test]
263 fn test_l2_squared_penalty() {
264 let lambda = 0.1f64;
265 let ar = ActivityRegularization::l2_squared(lambda);
266
267 let activations = Array1::from_vec(vec![1.0f64, 2.0, 3.0]);
268 let penalty = ar
269 .penalty(&activations)
270 .expect("ar.penalty succeeds in test_l2_squared_penalty");
271
272 assert_abs_diff_eq!(penalty, lambda * 14.0, epsilon = 1e-10);
274 }
275
276 #[test]
277 fn test_l1_gradients() {
278 let lambda = 0.1f64;
279 let ar = ActivityRegularization::l1(lambda);
280
281 let activations = Array1::from_vec(vec![1.0f64, -2.0, 0.0]);
282 let mut gradients = Array1::zeros(3);
283
284 let penalty = ar
285 .apply(&activations, &mut gradients)
286 .expect("apply succeeds in test_l1_gradients");
287
288 assert_abs_diff_eq!(gradients[0], lambda, epsilon = 1e-10); assert_abs_diff_eq!(gradients[1], -lambda, epsilon = 1e-10); assert_abs_diff_eq!(gradients[2], 0.0, epsilon = 1e-10); assert_abs_diff_eq!(penalty, lambda * 3.0, epsilon = 1e-10);
295 }
296
297 #[test]
298 fn test_l2_gradients() {
299 let lambda = 0.1f64;
300 let ar = ActivityRegularization::l2(lambda);
301
302 let activations = Array1::from_vec(vec![3.0f64, 4.0]);
303 let mut gradients = Array1::zeros(2);
304
305 let penalty = ar
306 .apply(&activations, &mut gradients)
307 .expect("apply succeeds in test_l2_gradients");
308
309 assert_abs_diff_eq!(gradients[0], lambda * 3.0 / 5.0, epsilon = 1e-10);
312 assert_abs_diff_eq!(gradients[1], lambda * 4.0 / 5.0, epsilon = 1e-10);
313
314 assert_abs_diff_eq!(penalty, lambda * 5.0, epsilon = 1e-10);
316 }
317
318 #[test]
319 fn test_l2_gradients_zero_activations() {
320 let lambda = 0.1f64;
321 let ar = ActivityRegularization::l2(lambda);
322
323 let activations = Array1::from_vec(vec![0.0f64, 0.0]);
324 let mut gradients = Array1::zeros(2);
325
326 let penalty = ar
327 .apply(&activations, &mut gradients)
328 .expect("apply succeeds in test_l2_gradients_zero_activations");
329
330 assert_abs_diff_eq!(gradients[0], 0.0, epsilon = 1e-10);
332 assert_abs_diff_eq!(gradients[1], 0.0, epsilon = 1e-10);
333
334 assert_abs_diff_eq!(penalty, 0.0, epsilon = 1e-10);
336 }
337
338 #[test]
339 fn test_l2_squared_gradients() {
340 let lambda = 0.1f64;
341 let ar = ActivityRegularization::l2_squared(lambda);
342
343 let activations = Array1::from_vec(vec![2.0f64, 3.0]);
344 let mut gradients = Array1::zeros(2);
345
346 let penalty = ar
347 .apply(&activations, &mut gradients)
348 .expect("apply succeeds in test_l2_squared_gradients");
349
350 assert_abs_diff_eq!(gradients[0], lambda * 2.0 * 2.0, epsilon = 1e-10);
352 assert_abs_diff_eq!(gradients[1], lambda * 2.0 * 3.0, epsilon = 1e-10);
353
354 assert_abs_diff_eq!(penalty, lambda * 13.0, epsilon = 1e-10);
356 }
357
358 #[test]
359 fn test_2d_activations() {
360 let lambda = 0.1f64;
361 let ar = ActivityRegularization::l1(lambda);
362
363 let activations = Array2::from_shape_vec((2, 2), vec![1.0f64, 2.0, -3.0, 4.0])
364 .expect("Array2::from_shape_vec succeeds in test_2d_activations");
365 let penalty = ar
366 .penalty(&activations)
367 .expect("ar.penalty succeeds in test_2d_activations");
368
369 assert_abs_diff_eq!(penalty, lambda * 10.0, epsilon = 1e-10);
371 }
372
373 #[test]
374 fn test_regularizer_trait() {
375 let lambda = 0.1f64;
376 let ar = ActivityRegularization::l1(lambda);
377
378 let activations = array![1.0f64, 2.0, 3.0];
379 let mut gradients = Array1::zeros(3);
380
381 let penalty1 = ar
383 .penalty(&activations)
384 .expect("ar.penalty succeeds in test_regularizer_trait");
385 let penalty2 = ar
386 .apply(&activations, &mut gradients)
387 .expect("apply succeeds in test_regularizer_trait");
388
389 assert_abs_diff_eq!(penalty1, penalty2, epsilon = 1e-10);
390
391 assert_abs_diff_eq!(gradients[0], lambda, epsilon = 1e-10);
393 assert_abs_diff_eq!(gradients[1], lambda, epsilon = 1e-10);
394 assert_abs_diff_eq!(gradients[2], lambda, epsilon = 1e-10);
395 }
396}