1use crate::NeuralResult;
15use scirs2_core::ndarray::{Array1, Array2, ScalarOperand};
16use scirs2_core::random::{thread_rng, Normal};
17use sklears_core::{error::SklearsError, types::FloatBounds};
18
19#[cfg(feature = "serde")]
20use serde::{Deserialize, Serialize};
21
22#[derive(Debug, Clone, Copy, PartialEq)]
24#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
25pub enum EBMTrainingAlgorithm {
26 ContrastiveDivergence {
28 k_steps: usize,
30 },
31 PersistentCD {
33 k_steps: usize,
35 },
36 ScoreMatching,
38 DenoisingScoreMatching {
40 noise_std: f64,
42 },
43 MaximumLikelihoodMCMC {
45 mcmc_steps: usize,
47 },
48}
49
50#[derive(Debug, Clone, Copy, PartialEq)]
52#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
53pub enum SamplingMethod {
54 Gibbs,
56 Langevin {
58 step_size: f64,
60 num_steps: usize,
62 },
63 HMC {
65 step_size: f64,
67 num_leapfrog: usize,
69 },
70}
71
72#[derive(Debug, Clone)]
74#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
75pub struct EBMConfig {
76 pub input_dim: usize,
78 pub hidden_dims: Vec<usize>,
80 pub training_algorithm: EBMTrainingAlgorithm,
82 pub sampling_method: SamplingMethod,
84 pub learning_rate: f64,
86 pub n_iterations: usize,
88 pub batch_size: usize,
90 pub use_bias: bool,
92}
93
94impl Default for EBMConfig {
95 fn default() -> Self {
96 Self {
97 input_dim: 784,
98 hidden_dims: vec![512, 256],
99 training_algorithm: EBMTrainingAlgorithm::ContrastiveDivergence { k_steps: 1 },
100 sampling_method: SamplingMethod::Langevin {
101 step_size: 0.01,
102 num_steps: 100,
103 },
104 learning_rate: 0.001,
105 n_iterations: 1000,
106 batch_size: 128,
107 use_bias: true,
108 }
109 }
110}
111
112#[derive(Debug)]
114#[allow(dead_code)] pub struct EnergyNetwork<T: FloatBounds> {
116 weights: Vec<Array2<T>>,
118 biases: Vec<Array1<T>>,
120 input_dim: usize,
122 cached_activations: Vec<Array2<T>>,
124}
125
126impl<T: FloatBounds + ScalarOperand> EnergyNetwork<T> {
127 pub fn new(input_dim: usize, hidden_dims: Vec<usize>) -> Self {
129 let mut rng = thread_rng();
130 let mut weights = Vec::new();
131 let mut biases = Vec::new();
132
133 let mut prev_dim = input_dim;
134 for &hidden_dim in &hidden_dims {
135 let std = (2.0 / prev_dim as f64).sqrt();
136 let w = Array2::from_shape_fn((prev_dim, hidden_dim), |_| {
137 T::from(
138 rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params"))
139 * std,
140 )
141 .unwrap_or_else(|| T::zero())
142 });
143 let b = Array1::zeros(hidden_dim);
144 weights.push(w);
145 biases.push(b);
146 prev_dim = hidden_dim;
147 }
148
149 let w = Array2::from_shape_fn((prev_dim, 1), |_| {
151 T::from(
152 rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params"))
153 * 0.01,
154 )
155 .unwrap_or_else(|| T::zero())
156 });
157 let b = Array1::zeros(1);
158 weights.push(w);
159 biases.push(b);
160
161 Self {
162 weights,
163 biases,
164 input_dim,
165 cached_activations: Vec::new(),
166 }
167 }
168
169 pub fn energy(&mut self, x: &Array2<T>) -> NeuralResult<Array1<T>> {
171 self.cached_activations.clear();
172 let mut h = x.clone();
173 self.cached_activations.push(h.clone());
174
175 for (i, (w, b)) in self.weights.iter().zip(self.biases.iter()).enumerate() {
177 h = h.dot(w);
178 for j in 0..h.nrows() {
179 h.row_mut(j).scaled_add(T::one(), &b.view());
180 }
181
182 if i < self.weights.len() - 1 {
184 h.mapv_inplace(|x| if x > T::zero() { x } else { T::zero() });
185 }
186
187 self.cached_activations.push(h.clone());
188 }
189
190 Ok(h.column(0).to_owned())
192 }
193
194 pub fn energy_gradient(&mut self, x: &Array2<T>) -> NeuralResult<Array2<T>> {
196 let _ = self.energy(x)?;
198
199 let batch_size = x.nrows();
200 let mut grad = Array2::ones((batch_size, 1));
201
202 for i in (0..self.weights.len()).rev() {
204 let h = &self.cached_activations[i + 1];
205
206 if i < self.weights.len() - 1 {
208 let activation_grad =
210 h.mapv(|xi| if xi > T::zero() { T::one() } else { T::zero() });
211 for j in 0..grad.nrows() {
212 for k in 0..grad.ncols() {
213 grad[[j, k]] *= activation_grad[[j, k]];
214 }
215 }
216 }
217
218 grad = grad.dot(&self.weights[i].t());
220 }
221
222 Ok(grad)
223 }
224
225 pub fn num_parameters(&self) -> usize {
227 self.weights.iter().map(|w| w.len()).sum::<usize>()
228 + self.biases.iter().map(|b| b.len()).sum::<usize>()
229 }
230}
231
232pub struct EnergyBasedModel<T: FloatBounds> {
234 energy_net: EnergyNetwork<T>,
236 config: EBMConfig,
238 persistent_chain: Option<Array2<T>>,
240}
241
242impl<T: FloatBounds + ScalarOperand> EnergyBasedModel<T> {
243 pub fn new(config: EBMConfig) -> Self {
245 let energy_net = EnergyNetwork::new(config.input_dim, config.hidden_dims.clone());
246
247 Self {
248 energy_net,
249 config,
250 persistent_chain: None,
251 }
252 }
253
254 pub fn energy(&mut self, x: &Array2<T>) -> NeuralResult<Array1<T>> {
256 self.energy_net.energy(x)
257 }
258
259 pub fn sample(&mut self, n_samples: usize) -> NeuralResult<Array2<T>> {
261 match self.config.sampling_method {
262 SamplingMethod::Gibbs => self.sample_gibbs(n_samples),
263 SamplingMethod::Langevin {
264 step_size,
265 num_steps,
266 } => self.sample_langevin(n_samples, step_size, num_steps),
267 SamplingMethod::HMC {
268 step_size,
269 num_leapfrog,
270 } => self.sample_hmc(n_samples, step_size, num_leapfrog),
271 }
272 }
273
274 fn sample_gibbs(&mut self, n_samples: usize) -> NeuralResult<Array2<T>> {
276 let mut rng = thread_rng();
277
278 let mut samples = Array2::from_shape_fn((n_samples, self.config.input_dim), |_| {
280 if rng.random::<f64>() < 0.5 {
281 T::zero()
282 } else {
283 T::one()
284 }
285 });
286
287 let num_iterations = 100;
289 for _ in 0..num_iterations {
290 for i in 0..self.config.input_dim {
291 let mut x_0 = samples.clone();
293 for j in 0..n_samples {
294 x_0[[j, i]] = T::zero();
295 }
296 let energy_0 = self.energy(&x_0)?;
297
298 let mut x_1 = samples.clone();
300 for j in 0..n_samples {
301 x_1[[j, i]] = T::one();
302 }
303 let energy_1 = self.energy(&x_1)?;
304
305 for j in 0..n_samples {
307 let prob_1 = T::one() / (T::one() + (energy_1[j] - energy_0[j]).exp());
308 samples[[j, i]] = if rng.random::<f64>() < prob_1.to_f64().unwrap_or(0.0) {
309 T::one()
310 } else {
311 T::zero()
312 };
313 }
314 }
315 }
316
317 Ok(samples)
318 }
319
320 fn sample_langevin(
322 &mut self,
323 n_samples: usize,
324 step_size: f64,
325 num_steps: usize,
326 ) -> NeuralResult<Array2<T>> {
327 let mut rng = thread_rng();
328
329 let mut samples = Array2::from_shape_fn((n_samples, self.config.input_dim), |_| {
331 T::from(rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")))
332 .unwrap_or_else(|| T::zero())
333 });
334
335 let step_size_t = T::from(step_size).unwrap_or_else(|| T::zero());
336 let noise_scale = T::from((2.0 * step_size).sqrt()).unwrap_or_else(|| T::zero());
337
338 for _ in 0..num_steps {
340 let grad = self.energy_net.energy_gradient(&samples)?;
342
343 let noise = Array2::from_shape_fn(samples.dim(), |_| {
345 T::from(
346 rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")),
347 )
348 .unwrap_or_else(|| T::zero())
349 });
350
351 samples = &samples - &grad.mapv(|g| g * step_size_t) + &noise.mapv(|n| n * noise_scale);
353 }
354
355 Ok(samples)
356 }
357
358 fn sample_hmc(
360 &mut self,
361 n_samples: usize,
362 step_size: f64,
363 num_leapfrog: usize,
364 ) -> NeuralResult<Array2<T>> {
365 let mut rng = thread_rng();
366
367 let mut q = Array2::from_shape_fn((n_samples, self.config.input_dim), |_| {
369 T::from(rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")))
370 .unwrap_or_else(|| T::zero())
371 });
372
373 let num_iterations = 100;
374 let step_size_t = T::from(step_size).unwrap_or_else(|| T::zero());
375
376 for _ in 0..num_iterations {
377 let mut p = Array2::from_shape_fn(q.dim(), |_| {
379 T::from(
380 rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")),
381 )
382 .unwrap_or_else(|| T::zero())
383 });
384
385 let q_old = q.clone();
386 let p_old = p.clone();
387
388 let grad = self.energy_net.energy_gradient(&q)?;
390 let p_half =
391 &p - &grad.mapv(|g| g * step_size_t * T::from(0.5).unwrap_or_else(|| T::zero()));
392
393 for i in 0..num_leapfrog {
395 q = &q + &p_half.mapv(|pi| pi * step_size_t);
397
398 if i < num_leapfrog - 1 {
400 let grad = self.energy_net.energy_gradient(&q)?;
401 let _p_updated = &p_half - &grad.mapv(|g| g * step_size_t);
403 let _ = &p; }
405 }
406
407 let grad = self.energy_net.energy_gradient(&q)?;
409 p = &p_half
410 - &grad.mapv(|g| g * step_size_t * T::from(0.5).unwrap_or_else(|| T::zero()));
411
412 let energy_old = self.energy(&q_old)?;
414 let energy_new = self.energy(&q)?;
415
416 let kinetic_old = p_old.mapv(|pi| pi * pi).sum();
417 let kinetic_new = p.mapv(|pi| pi * pi).sum();
418
419 for j in 0..n_samples {
420 let h_old = energy_old[j] + kinetic_old / T::from(2.0).unwrap_or_else(|| T::zero());
421 let h_new = energy_new[j] + kinetic_new / T::from(2.0).unwrap_or_else(|| T::zero());
422
423 let accept_prob = (-(h_new - h_old)).exp();
424
425 if rng.random::<f64>() > accept_prob.to_f64().unwrap_or(0.0) {
426 for k in 0..self.config.input_dim {
428 q[[j, k]] = q_old[[j, k]];
429 }
430 }
431 }
432 }
433
434 Ok(q)
435 }
436
437 pub fn train_contrastive_divergence(
439 &mut self,
440 x: &Array2<T>,
441 k_steps: usize,
442 ) -> NeuralResult<T> {
443 let _batch_size = x.nrows();
444
445 let energy_pos = self.energy(x)?;
447 let _grad_pos = self.energy_net.energy_gradient(x)?;
448
449 let x_neg = match &self.persistent_chain {
451 Some(chain)
452 if matches!(
453 self.config.training_algorithm,
454 EBMTrainingAlgorithm::PersistentCD { .. }
455 ) =>
456 {
457 let mut chain = chain.clone();
459 for _ in 0..k_steps {
460 let grad = self.energy_net.energy_gradient(&chain)?;
461 let noise = Array2::from_shape_fn(chain.dim(), |_| {
462 T::from(thread_rng().sample::<f64, _>(
463 Normal::new(0.0, 1.0).expect("valid distribution params"),
464 ))
465 .expect("value should be present")
466 });
467 chain = &chain - &grad.mapv(|g| g * T::from(0.01).unwrap_or_else(|| T::zero()))
468 + &noise.mapv(|n| n * T::from(0.1).unwrap_or_else(|| T::zero()));
469 }
470 self.persistent_chain = Some(chain.clone());
471 chain
472 }
473 _ => {
474 let mut x_neg = x.clone();
476 for _ in 0..k_steps {
477 let grad = self.energy_net.energy_gradient(&x_neg)?;
478 let noise = Array2::from_shape_fn(x_neg.dim(), |_| {
479 T::from(thread_rng().sample::<f64, _>(
480 Normal::new(0.0, 1.0).expect("valid distribution params"),
481 ))
482 .expect("value should be present")
483 });
484 x_neg = &x_neg - &grad.mapv(|g| g * T::from(0.01).unwrap_or_else(|| T::zero()))
485 + &noise.mapv(|n| n * T::from(0.1).unwrap_or_else(|| T::zero()));
486 }
487 x_neg
488 }
489 };
490
491 let energy_neg = self.energy(&x_neg)?;
492
493 let loss = (energy_pos
495 .mean()
496 .expect("mean should not fail on non-empty array")
497 - energy_neg
498 .mean()
499 .expect("mean should not fail on non-empty array"))
500 .abs();
501
502 Ok(loss)
503 }
504
505 pub fn config(&self) -> &EBMConfig {
507 &self.config
508 }
509
510 pub fn num_parameters(&self) -> usize {
512 self.energy_net.num_parameters()
513 }
514}
515
516#[derive(Debug)]
518pub struct HopfieldNetwork<T: FloatBounds> {
519 weights: Array2<T>,
521 n_units: usize,
523 patterns: Vec<Array1<T>>,
525}
526
527impl<T: FloatBounds + ScalarOperand> HopfieldNetwork<T> {
528 pub fn new(n_units: usize) -> Self {
530 let weights = Array2::zeros((n_units, n_units));
531
532 Self {
533 weights,
534 n_units,
535 patterns: Vec::new(),
536 }
537 }
538
539 pub fn store_pattern(&mut self, pattern: &Array1<T>) -> NeuralResult<()> {
541 if pattern.len() != self.n_units {
542 return Err(SklearsError::InvalidParameter {
543 name: "pattern".to_string(),
544 reason: format!(
545 "Pattern length {} does not match network size {}",
546 pattern.len(),
547 self.n_units
548 ),
549 });
550 }
551
552 self.patterns.push(pattern.clone());
554
555 for i in 0..self.n_units {
557 for j in 0..self.n_units {
558 if i != j {
559 self.weights[[i, j]] += pattern[i] * pattern[j];
560 }
561 }
562 }
563
564 let n_patterns = T::from(self.patterns.len() as f64).unwrap_or_else(|| T::zero());
566 self.weights.mapv_inplace(|w| w / n_patterns);
567
568 Ok(())
569 }
570
571 pub fn recall(&self, initial: &Array1<T>, max_iterations: usize) -> NeuralResult<Array1<T>> {
573 if initial.len() != self.n_units {
574 return Err(SklearsError::InvalidParameter {
575 name: "initial".to_string(),
576 reason: format!(
577 "Initial state length {} does not match network size {}",
578 initial.len(),
579 self.n_units
580 ),
581 });
582 }
583
584 let mut state = initial.clone();
585
586 for _ in 0..max_iterations {
588 let mut new_state = state.clone();
589
590 for i in 0..self.n_units {
592 let activation = self.weights.row(i).dot(&state);
593 new_state[i] = if activation >= T::zero() {
594 T::one()
595 } else {
596 -T::one()
597 };
598 }
599
600 if new_state == state {
602 break;
603 }
604
605 state = new_state;
606 }
607
608 Ok(state)
609 }
610
611 pub fn energy(&self, state: &Array1<T>) -> T {
613 let mut energy = T::zero();
614
615 for i in 0..self.n_units {
616 for j in 0..self.n_units {
617 energy -= self.weights[[i, j]] * state[i] * state[j];
618 }
619 }
620
621 energy / T::from(2.0).unwrap_or_else(|| T::zero())
622 }
623
624 pub fn num_patterns(&self) -> usize {
626 self.patterns.len()
627 }
628}
629
630#[cfg(test)]
631mod tests {
632 use super::*;
633
634 #[test]
635 fn test_energy_network_creation() {
636 let network: EnergyNetwork<f64> = EnergyNetwork::new(10, vec![32, 16]);
637 assert_eq!(network.input_dim, 10);
638 assert!(network.num_parameters() > 0);
639 }
640
641 #[test]
642 fn test_energy_computation() {
643 let mut network: EnergyNetwork<f64> = EnergyNetwork::new(8, vec![16]);
644 let x = Array2::from_shape_fn((4, 8), |(i, j)| (i + j) as f64 * 0.1);
645
646 let energy = network.energy(&x).expect("operation should succeed");
647 assert_eq!(energy.len(), 4);
648 assert!(energy.iter().all(|&e| e.is_finite()));
649 }
650
651 #[test]
652 fn test_energy_gradient() {
653 let mut network: EnergyNetwork<f64> = EnergyNetwork::new(6, vec![12]);
654 let x = Array2::from_shape_fn((3, 6), |(i, j)| (i + j) as f64 * 0.1);
655
656 let grad = network
657 .energy_gradient(&x)
658 .expect("operation should succeed");
659 assert_eq!(grad.dim(), x.dim());
660 assert!(grad.iter().all(|&g| g.is_finite()));
661 }
662
663 #[test]
664 fn test_ebm_creation() {
665 let config = EBMConfig {
666 input_dim: 10,
667 hidden_dims: vec![32],
668 ..Default::default()
669 };
670
671 let ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
672 assert!(ebm.num_parameters() > 0);
673 }
674
675 #[test]
676 fn test_langevin_sampling() {
677 let config = EBMConfig {
678 input_dim: 8,
679 hidden_dims: vec![16],
680 sampling_method: SamplingMethod::Langevin {
681 step_size: 0.01,
682 num_steps: 10,
683 },
684 ..Default::default()
685 };
686
687 let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
688 let samples = ebm.sample(5).expect("sampling should succeed");
689
690 assert_eq!(samples.nrows(), 5);
691 assert_eq!(samples.ncols(), 8);
692 }
693
694 #[test]
695 fn test_gibbs_sampling() {
696 let config = EBMConfig {
697 input_dim: 6,
698 hidden_dims: vec![12],
699 sampling_method: SamplingMethod::Gibbs,
700 ..Default::default()
701 };
702
703 let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
704 let samples = ebm.sample(4).expect("sampling should succeed");
705
706 assert_eq!(samples.nrows(), 4);
707 assert_eq!(samples.ncols(), 6);
708 assert!(samples.iter().all(|&x| x == 0.0 || x == 1.0));
710 }
711
712 #[test]
713 fn test_contrastive_divergence() {
714 let config = EBMConfig {
715 input_dim: 8,
716 hidden_dims: vec![16],
717 ..Default::default()
718 };
719
720 let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
721 let x = Array2::from_shape_fn((4, 8), |(i, j)| (i + j) as f64 * 0.1);
722
723 let loss = ebm
724 .train_contrastive_divergence(&x, 1)
725 .expect("operation should succeed");
726 assert!(loss.is_finite());
727 assert!(loss >= 0.0);
728 }
729
730 #[test]
731 fn test_hopfield_network_creation() {
732 let network: HopfieldNetwork<f64> = HopfieldNetwork::new(10);
733 assert_eq!(network.n_units, 10);
734 assert_eq!(network.num_patterns(), 0);
735 }
736
737 #[test]
738 fn test_hopfield_store_pattern() {
739 let mut network: HopfieldNetwork<f64> = HopfieldNetwork::new(5);
740 let pattern = Array1::from_vec(vec![1.0, -1.0, 1.0, -1.0, 1.0]);
741
742 network
743 .store_pattern(&pattern)
744 .expect("operation should succeed");
745 assert_eq!(network.num_patterns(), 1);
746 }
747
748 #[test]
749 fn test_hopfield_recall() {
750 let mut network: HopfieldNetwork<f64> = HopfieldNetwork::new(4);
751
752 let pattern = Array1::from_vec(vec![1.0, 1.0, -1.0, -1.0]);
754 network
755 .store_pattern(&pattern)
756 .expect("operation should succeed");
757
758 let noisy = Array1::from_vec(vec![1.0, -1.0, -1.0, -1.0]);
760 let recalled = network
761 .recall(&noisy, 10)
762 .expect("operation should succeed");
763
764 assert_eq!(recalled.len(), 4);
765 assert!(recalled.iter().all(|&x| x == 1.0 || x == -1.0));
767 }
768
769 #[test]
770 fn test_hopfield_energy() {
771 let mut network: HopfieldNetwork<f64> = HopfieldNetwork::new(4);
772 let pattern = Array1::from_vec(vec![1.0, -1.0, 1.0, -1.0]);
773 network
774 .store_pattern(&pattern)
775 .expect("operation should succeed");
776
777 let energy = network.energy(&pattern);
778 assert!(energy.is_finite());
779
780 let random = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0]);
782 let energy_random = network.energy(&random);
783
784 assert!(energy < energy_random);
785 }
786
787 #[test]
788 fn test_hmc_sampling() {
789 let config = EBMConfig {
790 input_dim: 6,
791 hidden_dims: vec![12],
792 sampling_method: SamplingMethod::HMC {
793 step_size: 0.01,
794 num_leapfrog: 5,
795 },
796 ..Default::default()
797 };
798
799 let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
800 let samples = ebm.sample(3).expect("sampling should succeed");
801
802 assert_eq!(samples.nrows(), 3);
803 assert_eq!(samples.ncols(), 6);
804 assert!(samples.iter().all(|&x| x.is_finite()));
805 }
806}