optirs_core/regularizers/
entropy.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)]
25pub enum EntropyRegularizerType {
26 MaximizeEntropy,
28 MinimizeEntropy,
30}
31
32#[derive(Debug, Clone, Copy)]
40pub struct EntropyRegularization<A: Float + FromPrimitive + Debug> {
41 pub lambda: A,
43 pub epsilon: A,
45 pub reg_type: EntropyRegularizerType,
47}
48
49impl<A: Float + FromPrimitive + Debug + Send + Sync> EntropyRegularization<A> {
50 pub fn new(lambda: A, regtype: EntropyRegularizerType) -> Self {
61 let epsilon =
62 A::from_f64(1e-8).expect("EntropyRegularization: default epsilon (1e-8) must fit in A");
63 Self {
64 lambda,
65 epsilon,
66 reg_type: regtype,
67 }
68 }
69
70 pub fn new_with_epsilon(lambda: A, epsilon: A, regtype: EntropyRegularizerType) -> Self {
82 Self {
83 lambda,
84 epsilon,
85 reg_type: regtype,
86 }
87 }
88
89 pub fn calculate_entropy<S, D>(&self, probs: &ArrayBase<S, D>) -> A
99 where
100 S: Data<Elem = A>,
101 D: Dimension,
102 {
103 let safe_probs = probs.mapv(|p| {
105 if p < self.epsilon {
106 self.epsilon
107 } else if p > (A::one() - self.epsilon) {
108 A::one() - self.epsilon
109 } else {
110 p
111 }
112 });
113
114 let neg_entropy = safe_probs.mapv(|p| p * p.ln()).sum();
116 -neg_entropy
117 }
118
119 fn entropy_gradient<S, D>(&self, probs: &ArrayBase<S, D>) -> Array<A, D>
129 where
130 S: Data<Elem = A>,
131 D: Dimension,
132 {
133 let safe_probs = probs.mapv(|p| {
135 if p < self.epsilon {
136 self.epsilon
137 } else if p > (A::one() - self.epsilon) {
138 A::one() - self.epsilon
139 } else {
140 p
141 }
142 });
143
144 let gradient = safe_probs.mapv(|p| -(A::one() + p.ln()));
146
147 match self.reg_type {
149 EntropyRegularizerType::MaximizeEntropy => gradient,
150 EntropyRegularizerType::MinimizeEntropy => gradient.mapv(|g| -g),
151 }
152 }
153}
154
155impl<A, D> Regularizer<A, D> for EntropyRegularization<A>
156where
157 A: Float + ScalarOperand + Debug + FromPrimitive + Send + Sync,
158 D: Dimension,
159{
160 fn apply(&self, params: &Array<A, D>, gradients: &mut Array<A, D>) -> Result<A> {
161 let entropy = self.calculate_entropy(params);
163
164 let entropy_grads = self.entropy_gradient(params);
166
167 gradients.zip_mut_with(&entropy_grads, |g, &e| *g = *g + self.lambda * e);
169
170 let penalty = match self.reg_type {
174 EntropyRegularizerType::MaximizeEntropy => -self.lambda * entropy,
175 EntropyRegularizerType::MinimizeEntropy => self.lambda * entropy,
176 };
177
178 Ok(penalty)
179 }
180
181 fn penalty(&self, params: &Array<A, D>) -> Result<A> {
182 let entropy = self.calculate_entropy(params);
184
185 let penalty = match self.reg_type {
188 EntropyRegularizerType::MaximizeEntropy => -self.lambda * entropy,
189 EntropyRegularizerType::MinimizeEntropy => self.lambda * entropy,
190 };
191
192 Ok(penalty)
193 }
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199 use approx::assert_abs_diff_eq;
200 use scirs2_core::ndarray::Array1;
201
202 #[test]
203 fn test_entropy_regularization_creation() {
204 let er = EntropyRegularization::new(0.1f64, EntropyRegularizerType::MaximizeEntropy);
205 assert_eq!(er.lambda, 0.1);
206 assert_eq!(er.epsilon, 1e-8);
207 match er.reg_type {
208 EntropyRegularizerType::MaximizeEntropy => (),
209 _ => panic!("Wrong regularizer type"),
210 }
211
212 let er = EntropyRegularization::new_with_epsilon(
213 0.2f64,
214 1e-10,
215 EntropyRegularizerType::MinimizeEntropy,
216 );
217 assert_eq!(er.lambda, 0.2);
218 assert_eq!(er.epsilon, 1e-10);
219 match er.reg_type {
220 EntropyRegularizerType::MinimizeEntropy => (),
221 _ => panic!("Wrong regularizer type"),
222 }
223 }
224
225 #[test]
226 fn test_calculate_entropy() {
227 let uniform = Array1::from_vec(vec![0.25f64, 0.25, 0.25, 0.25]);
229 let er = EntropyRegularization::new(1.0f64, EntropyRegularizerType::MaximizeEntropy);
230 let entropy = er.calculate_entropy(&uniform);
231
232 let expected = (4.0f64).ln();
234 assert_abs_diff_eq!(entropy, expected, epsilon = 1e-6);
235
236 let peaked = Array1::from_vec(vec![0.01f64, 0.01, 0.97, 0.01]);
238 let entropy = er.calculate_entropy(&peaked);
239 assert!(entropy < expected); }
241
242 #[test]
243 fn test_entropy_gradient() {
244 let er = EntropyRegularization::new(1.0f64, EntropyRegularizerType::MaximizeEntropy);
245
246 let uniform = Array1::from_vec(vec![0.25f64, 0.25, 0.25, 0.25]);
248 let grads = er.entropy_gradient(&uniform);
249
250 let expected = -(1.0 + 0.25f64.ln());
252 for &g in grads.iter() {
253 assert_abs_diff_eq!(g, expected, epsilon = 1e-6);
254 }
255
256 let peaked = Array1::from_vec(vec![0.1f64, 0.1, 0.7, 0.1]);
258 let grads = er.entropy_gradient(&peaked);
259
260 assert!(grads[2].abs() < grads[0].abs());
264 }
265
266 #[test]
267 fn test_maximize_entropy_penalty() {
268 let er = EntropyRegularization::new(1.0f64, EntropyRegularizerType::MaximizeEntropy);
270
271 let uniform = Array1::from_vec(vec![0.25f64, 0.25, 0.25, 0.25]);
273 let penalty = er
274 .penalty(&uniform)
275 .expect("er.penalty succeeds in test_maximize_entropy_penalty");
276
277 let peaked = Array1::from_vec(vec![0.01f64, 0.01, 0.97, 0.01]);
279 let peaked_penalty = er
280 .penalty(&peaked)
281 .expect("er.penalty succeeds in test_maximize_entropy_penalty");
282
283 assert!(peaked_penalty > penalty);
286 }
287
288 #[test]
289 fn test_minimize_entropy_penalty() {
290 let er = EntropyRegularization::new(1.0f64, EntropyRegularizerType::MinimizeEntropy);
292
293 let uniform = Array1::from_vec(vec![0.25f64, 0.25, 0.25, 0.25]);
295 let penalty = er
296 .penalty(&uniform)
297 .expect("er.penalty succeeds in test_minimize_entropy_penalty");
298
299 let peaked = Array1::from_vec(vec![0.01f64, 0.01, 0.97, 0.01]);
301 let peaked_penalty = er
302 .penalty(&peaked)
303 .expect("er.penalty succeeds in test_minimize_entropy_penalty");
304
305 assert!(penalty > peaked_penalty);
308 }
309
310 #[test]
311 fn test_apply_gradients() {
312 let lambda = 0.5f64;
313 let er = EntropyRegularization::new(lambda, EntropyRegularizerType::MaximizeEntropy);
314
315 let probs = Array1::from_vec(vec![0.25f64, 0.25, 0.25, 0.25]);
316 let mut gradients = Array1::zeros(4);
317
318 let penalty = er
319 .apply(&probs, &mut gradients)
320 .expect("er.apply succeeds in test_apply_gradients");
321
322 assert!(gradients.iter().all(|&g| g != 0.0));
324
325 let first = gradients[0];
327 assert!(gradients.iter().all(|&g| (g - first).abs() < 1e-6));
328
329 let expected_grad = -lambda * (1.0 + 0.25f64.ln());
331 assert_abs_diff_eq!(gradients[0], expected_grad, epsilon = 1e-6);
332
333 let entropy = (4.0f64).ln(); let expected_penalty = -lambda * entropy; assert_abs_diff_eq!(penalty, expected_penalty, epsilon = 1e-6);
337 }
338
339 #[test]
340 fn test_regularizer_trait() {
341 let er = EntropyRegularization::new(0.1f64, EntropyRegularizerType::MaximizeEntropy);
343
344 let probs = Array1::from_vec(vec![0.25f64, 0.25, 0.25, 0.25]);
345 let mut gradients = Array1::zeros(4);
346
347 let penalty1 = er
349 .apply(&probs, &mut gradients)
350 .expect("er.apply succeeds in test_regularizer_trait");
351 let penalty2 = er
352 .penalty(&probs)
353 .expect("er.penalty succeeds in test_regularizer_trait");
354
355 assert_abs_diff_eq!(penalty1, penalty2, epsilon = 1e-10);
356 }
357}