Skip to main content

sklears_neural/
ebm.rs

1//! Energy-Based Models (EBMs) for generative modeling and density estimation.
2//!
3//! This module implements various energy-based model architectures including:
4//! - Basic Energy-Based Models with contrastive divergence
5//! - Hopfield Networks
6//! - Restricted Boltzmann Machines (extended from rbm.rs)
7//! - Deep Energy Models
8//! - Score Matching for EBM training
9//! - Langevin dynamics for sampling
10//!
11//! Energy-based models assign a scalar energy to each configuration of variables
12//! and learn by making observed data have lower energy than unobserved data.
13
14use 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/// Training algorithm for energy-based models
23#[derive(Debug, Clone, Copy, PartialEq)]
24#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
25pub enum EBMTrainingAlgorithm {
26    /// Contrastive Divergence with k steps
27    ContrastiveDivergence {
28        /// Number of Gibbs sampling steps per CD update
29        k_steps: usize,
30    },
31    /// Persistent Contrastive Divergence
32    PersistentCD {
33        /// Number of Gibbs sampling steps per PCD update
34        k_steps: usize,
35    },
36    /// Score Matching
37    ScoreMatching,
38    /// Denoising Score Matching
39    DenoisingScoreMatching {
40        /// Standard deviation of the Gaussian noise added to inputs during training
41        noise_std: f64,
42    },
43    /// Maximum Likelihood with MCMC
44    MaximumLikelihoodMCMC {
45        /// Number of MCMC steps used to approximate the partition function gradient
46        mcmc_steps: usize,
47    },
48}
49
50/// Sampling method for EBM
51#[derive(Debug, Clone, Copy, PartialEq)]
52#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
53pub enum SamplingMethod {
54    /// Gibbs sampling
55    Gibbs,
56    /// Langevin dynamics with specific step size
57    Langevin {
58        /// Discretization step size for the Langevin update
59        step_size: f64,
60        /// Number of Langevin steps per sample
61        num_steps: usize,
62    },
63    /// Hamiltonian Monte Carlo
64    HMC {
65        /// Discretization step size (epsilon) for the leapfrog integrator
66        step_size: f64,
67        /// Number of leapfrog steps per HMC proposal
68        num_leapfrog: usize,
69    },
70}
71
72/// Configuration for energy-based model
73#[derive(Debug, Clone)]
74#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
75pub struct EBMConfig {
76    /// Input dimension
77    pub input_dim: usize,
78    /// Hidden layer sizes for energy network
79    pub hidden_dims: Vec<usize>,
80    /// Training algorithm
81    pub training_algorithm: EBMTrainingAlgorithm,
82    /// Sampling method
83    pub sampling_method: SamplingMethod,
84    /// Learning rate
85    pub learning_rate: f64,
86    /// Number of training iterations
87    pub n_iterations: usize,
88    /// Batch size
89    pub batch_size: usize,
90    /// Whether to use bias in energy network
91    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/// Energy network - maps input to scalar energy value
113#[derive(Debug)]
114#[allow(dead_code)] // input_dim retained for shape validation and future serialization
115pub struct EnergyNetwork<T: FloatBounds> {
116    /// Network weights
117    weights: Vec<Array2<T>>,
118    /// Network biases
119    biases: Vec<Array1<T>>,
120    /// Input dimension
121    input_dim: usize,
122    /// Cached activations for backprop
123    cached_activations: Vec<Array2<T>>,
124}
125
126impl<T: FloatBounds + ScalarOperand> EnergyNetwork<T> {
127    /// Create a new energy network
128    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        // Output layer (maps to single energy value)
150        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    /// Compute energy for input batch (forward pass)
170    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        // Forward through network
176        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            // Apply ReLU activation for all but last layer
183            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        // Return energy values (last layer output)
191        Ok(h.column(0).to_owned())
192    }
193
194    /// Compute gradient of energy with respect to input
195    pub fn energy_gradient(&mut self, x: &Array2<T>) -> NeuralResult<Array2<T>> {
196        // Compute energy first (populates cache)
197        let _ = self.energy(x)?;
198
199        let batch_size = x.nrows();
200        let mut grad = Array2::ones((batch_size, 1));
201
202        // Backpropagate through network
203        for i in (0..self.weights.len()).rev() {
204            let h = &self.cached_activations[i + 1];
205
206            // Gradient through activation
207            if i < self.weights.len() - 1 {
208                // ReLU derivative
209                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            // Gradient through linear layer
219            grad = grad.dot(&self.weights[i].t());
220        }
221
222        Ok(grad)
223    }
224
225    /// Get number of parameters
226    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
232/// Energy-Based Model
233pub struct EnergyBasedModel<T: FloatBounds> {
234    /// Energy network
235    energy_net: EnergyNetwork<T>,
236    /// Configuration
237    config: EBMConfig,
238    /// Persistent chain states (for PCD)
239    persistent_chain: Option<Array2<T>>,
240}
241
242impl<T: FloatBounds + ScalarOperand> EnergyBasedModel<T> {
243    /// Create a new energy-based model
244    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    /// Compute energy for input batch
255    pub fn energy(&mut self, x: &Array2<T>) -> NeuralResult<Array1<T>> {
256        self.energy_net.energy(x)
257    }
258
259    /// Sample from the model using configured sampling method
260    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    /// Sample using Gibbs sampling (for binary variables)
275    fn sample_gibbs(&mut self, n_samples: usize) -> NeuralResult<Array2<T>> {
276        let mut rng = thread_rng();
277
278        // Initialize random samples
279        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        // Gibbs sampling iterations
288        let num_iterations = 100;
289        for _ in 0..num_iterations {
290            for i in 0..self.config.input_dim {
291                // Compute energy with feature i = 0
292                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                // Compute energy with feature i = 1
299                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                // Sample from conditional probability
306                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    /// Sample using Langevin dynamics
321    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        // Initialize from random noise
330        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        // Langevin dynamics iterations
339        for _ in 0..num_steps {
340            // Compute gradient of energy
341            let grad = self.energy_net.energy_gradient(&samples)?;
342
343            // Add noise
344            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            // Update: x_{t+1} = x_t - step_size * ∇E(x_t) + noise
352            samples = &samples - &grad.mapv(|g| g * step_size_t) + &noise.mapv(|n| n * noise_scale);
353        }
354
355        Ok(samples)
356    }
357
358    /// Sample using Hamiltonian Monte Carlo
359    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        // Initialize positions
368        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            // Sample momentum
378            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            // Compute initial energy
389            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            // Leapfrog integration
394            for i in 0..num_leapfrog {
395                // Full step for position
396                q = &q + &p_half.mapv(|pi| pi * step_size_t);
397
398                // Full step for momentum (except last)
399                if i < num_leapfrog - 1 {
400                    let grad = self.energy_net.energy_gradient(&q)?;
401                    // Update momentum; shadowed intentionally - full HMC step placeholder
402                    let _p_updated = &p_half - &grad.mapv(|g| g * step_size_t);
403                    let _ = &p; // Keep p in scope for final step
404                }
405            }
406
407            // Half step for momentum at end
408            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            // Metropolis-Hastings acceptance
413            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                    // Reject - restore old position
427                    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    /// Train the model using contrastive divergence
438    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        // Positive phase: compute energy gradient on data
446        let energy_pos = self.energy(x)?;
447        let _grad_pos = self.energy_net.energy_gradient(x)?;
448
449        // Negative phase: sample from model
450        let x_neg = match &self.persistent_chain {
451            Some(chain)
452                if matches!(
453                    self.config.training_algorithm,
454                    EBMTrainingAlgorithm::PersistentCD { .. }
455                ) =>
456            {
457                // Use persistent chain for PCD
458                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                // Initialize from data for CD
475                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        // Compute loss (difference in energies)
494        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    /// Get configuration
506    pub fn config(&self) -> &EBMConfig {
507        &self.config
508    }
509
510    /// Get number of parameters
511    pub fn num_parameters(&self) -> usize {
512        self.energy_net.num_parameters()
513    }
514}
515
516/// Hopfield Network for associative memory
517#[derive(Debug)]
518pub struct HopfieldNetwork<T: FloatBounds> {
519    /// Weight matrix (symmetric)
520    weights: Array2<T>,
521    /// Number of units
522    n_units: usize,
523    /// Stored patterns
524    patterns: Vec<Array1<T>>,
525}
526
527impl<T: FloatBounds + ScalarOperand> HopfieldNetwork<T> {
528    /// Create a new Hopfield network
529    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    /// Store a pattern using Hebbian learning
540    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        // Store pattern for reference
553        self.patterns.push(pattern.clone());
554
555        // Update weights using outer product
556        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        // Normalize weights
565        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    /// Recall a pattern from partial/noisy input
572    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        // Iterate until convergence or max iterations
587        for _ in 0..max_iterations {
588            let mut new_state = state.clone();
589
590            // Update each unit
591            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            // Check for convergence
601            if new_state == state {
602                break;
603            }
604
605            state = new_state;
606        }
607
608        Ok(state)
609    }
610
611    /// Compute energy of a state
612    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    /// Get number of stored patterns
625    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        // Gibbs sampling produces binary values
709        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        // Store a pattern
753        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        // Recall from noisy version
759        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        // Should recall close to original pattern
766        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        // Energy should be lower for stored pattern
781        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}