use crate::NeuralResult;
use scirs2_core::ndarray::{Array1, Array2, ScalarOperand};
use scirs2_core::random::{thread_rng, Normal};
use sklears_core::{error::SklearsError, types::FloatBounds};
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum EBMTrainingAlgorithm {
ContrastiveDivergence {
k_steps: usize,
},
PersistentCD {
k_steps: usize,
},
ScoreMatching,
DenoisingScoreMatching {
noise_std: f64,
},
MaximumLikelihoodMCMC {
mcmc_steps: usize,
},
}
#[derive(Debug, Clone, Copy, PartialEq)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub enum SamplingMethod {
Gibbs,
Langevin {
step_size: f64,
num_steps: usize,
},
HMC {
step_size: f64,
num_leapfrog: usize,
},
}
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(Serialize, Deserialize))]
pub struct EBMConfig {
pub input_dim: usize,
pub hidden_dims: Vec<usize>,
pub training_algorithm: EBMTrainingAlgorithm,
pub sampling_method: SamplingMethod,
pub learning_rate: f64,
pub n_iterations: usize,
pub batch_size: usize,
pub use_bias: bool,
}
impl Default for EBMConfig {
fn default() -> Self {
Self {
input_dim: 784,
hidden_dims: vec![512, 256],
training_algorithm: EBMTrainingAlgorithm::ContrastiveDivergence { k_steps: 1 },
sampling_method: SamplingMethod::Langevin {
step_size: 0.01,
num_steps: 100,
},
learning_rate: 0.001,
n_iterations: 1000,
batch_size: 128,
use_bias: true,
}
}
}
#[derive(Debug)]
#[allow(dead_code)] pub struct EnergyNetwork<T: FloatBounds> {
weights: Vec<Array2<T>>,
biases: Vec<Array1<T>>,
input_dim: usize,
cached_activations: Vec<Array2<T>>,
}
impl<T: FloatBounds + ScalarOperand> EnergyNetwork<T> {
pub fn new(input_dim: usize, hidden_dims: Vec<usize>) -> Self {
let mut rng = thread_rng();
let mut weights = Vec::new();
let mut biases = Vec::new();
let mut prev_dim = input_dim;
for &hidden_dim in &hidden_dims {
let std = (2.0 / prev_dim as f64).sqrt();
let w = Array2::from_shape_fn((prev_dim, hidden_dim), |_| {
T::from(
rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params"))
* std,
)
.unwrap_or_else(|| T::zero())
});
let b = Array1::zeros(hidden_dim);
weights.push(w);
biases.push(b);
prev_dim = hidden_dim;
}
let w = Array2::from_shape_fn((prev_dim, 1), |_| {
T::from(
rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params"))
* 0.01,
)
.unwrap_or_else(|| T::zero())
});
let b = Array1::zeros(1);
weights.push(w);
biases.push(b);
Self {
weights,
biases,
input_dim,
cached_activations: Vec::new(),
}
}
pub fn energy(&mut self, x: &Array2<T>) -> NeuralResult<Array1<T>> {
self.cached_activations.clear();
let mut h = x.clone();
self.cached_activations.push(h.clone());
for (i, (w, b)) in self.weights.iter().zip(self.biases.iter()).enumerate() {
h = h.dot(w);
for j in 0..h.nrows() {
h.row_mut(j).scaled_add(T::one(), &b.view());
}
if i < self.weights.len() - 1 {
h.mapv_inplace(|x| if x > T::zero() { x } else { T::zero() });
}
self.cached_activations.push(h.clone());
}
Ok(h.column(0).to_owned())
}
pub fn energy_gradient(&mut self, x: &Array2<T>) -> NeuralResult<Array2<T>> {
let _ = self.energy(x)?;
let batch_size = x.nrows();
let mut grad = Array2::ones((batch_size, 1));
for i in (0..self.weights.len()).rev() {
let h = &self.cached_activations[i + 1];
if i < self.weights.len() - 1 {
let activation_grad =
h.mapv(|xi| if xi > T::zero() { T::one() } else { T::zero() });
for j in 0..grad.nrows() {
for k in 0..grad.ncols() {
grad[[j, k]] *= activation_grad[[j, k]];
}
}
}
grad = grad.dot(&self.weights[i].t());
}
Ok(grad)
}
pub fn num_parameters(&self) -> usize {
self.weights.iter().map(|w| w.len()).sum::<usize>()
+ self.biases.iter().map(|b| b.len()).sum::<usize>()
}
}
pub struct EnergyBasedModel<T: FloatBounds> {
energy_net: EnergyNetwork<T>,
config: EBMConfig,
persistent_chain: Option<Array2<T>>,
}
impl<T: FloatBounds + ScalarOperand> EnergyBasedModel<T> {
pub fn new(config: EBMConfig) -> Self {
let energy_net = EnergyNetwork::new(config.input_dim, config.hidden_dims.clone());
Self {
energy_net,
config,
persistent_chain: None,
}
}
pub fn energy(&mut self, x: &Array2<T>) -> NeuralResult<Array1<T>> {
self.energy_net.energy(x)
}
pub fn sample(&mut self, n_samples: usize) -> NeuralResult<Array2<T>> {
match self.config.sampling_method {
SamplingMethod::Gibbs => self.sample_gibbs(n_samples),
SamplingMethod::Langevin {
step_size,
num_steps,
} => self.sample_langevin(n_samples, step_size, num_steps),
SamplingMethod::HMC {
step_size,
num_leapfrog,
} => self.sample_hmc(n_samples, step_size, num_leapfrog),
}
}
fn sample_gibbs(&mut self, n_samples: usize) -> NeuralResult<Array2<T>> {
let mut rng = thread_rng();
let mut samples = Array2::from_shape_fn((n_samples, self.config.input_dim), |_| {
if rng.random::<f64>() < 0.5 {
T::zero()
} else {
T::one()
}
});
let num_iterations = 100;
for _ in 0..num_iterations {
for i in 0..self.config.input_dim {
let mut x_0 = samples.clone();
for j in 0..n_samples {
x_0[[j, i]] = T::zero();
}
let energy_0 = self.energy(&x_0)?;
let mut x_1 = samples.clone();
for j in 0..n_samples {
x_1[[j, i]] = T::one();
}
let energy_1 = self.energy(&x_1)?;
for j in 0..n_samples {
let prob_1 = T::one() / (T::one() + (energy_1[j] - energy_0[j]).exp());
samples[[j, i]] = if rng.random::<f64>() < prob_1.to_f64().unwrap_or(0.0) {
T::one()
} else {
T::zero()
};
}
}
}
Ok(samples)
}
fn sample_langevin(
&mut self,
n_samples: usize,
step_size: f64,
num_steps: usize,
) -> NeuralResult<Array2<T>> {
let mut rng = thread_rng();
let mut samples = Array2::from_shape_fn((n_samples, self.config.input_dim), |_| {
T::from(rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")))
.unwrap_or_else(|| T::zero())
});
let step_size_t = T::from(step_size).unwrap_or_else(|| T::zero());
let noise_scale = T::from((2.0 * step_size).sqrt()).unwrap_or_else(|| T::zero());
for _ in 0..num_steps {
let grad = self.energy_net.energy_gradient(&samples)?;
let noise = Array2::from_shape_fn(samples.dim(), |_| {
T::from(
rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")),
)
.unwrap_or_else(|| T::zero())
});
samples = &samples - &grad.mapv(|g| g * step_size_t) + &noise.mapv(|n| n * noise_scale);
}
Ok(samples)
}
fn sample_hmc(
&mut self,
n_samples: usize,
step_size: f64,
num_leapfrog: usize,
) -> NeuralResult<Array2<T>> {
let mut rng = thread_rng();
let mut q = Array2::from_shape_fn((n_samples, self.config.input_dim), |_| {
T::from(rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")))
.unwrap_or_else(|| T::zero())
});
let num_iterations = 100;
let step_size_t = T::from(step_size).unwrap_or_else(|| T::zero());
for _ in 0..num_iterations {
let mut p = Array2::from_shape_fn(q.dim(), |_| {
T::from(
rng.sample::<f64, _>(Normal::new(0.0, 1.0).expect("valid distribution params")),
)
.unwrap_or_else(|| T::zero())
});
let q_old = q.clone();
let p_old = p.clone();
let grad = self.energy_net.energy_gradient(&q)?;
let p_half =
&p - &grad.mapv(|g| g * step_size_t * T::from(0.5).unwrap_or_else(|| T::zero()));
for i in 0..num_leapfrog {
q = &q + &p_half.mapv(|pi| pi * step_size_t);
if i < num_leapfrog - 1 {
let grad = self.energy_net.energy_gradient(&q)?;
let _p_updated = &p_half - &grad.mapv(|g| g * step_size_t);
let _ = &p; }
}
let grad = self.energy_net.energy_gradient(&q)?;
p = &p_half
- &grad.mapv(|g| g * step_size_t * T::from(0.5).unwrap_or_else(|| T::zero()));
let energy_old = self.energy(&q_old)?;
let energy_new = self.energy(&q)?;
let kinetic_old = p_old.mapv(|pi| pi * pi).sum();
let kinetic_new = p.mapv(|pi| pi * pi).sum();
for j in 0..n_samples {
let h_old = energy_old[j] + kinetic_old / T::from(2.0).unwrap_or_else(|| T::zero());
let h_new = energy_new[j] + kinetic_new / T::from(2.0).unwrap_or_else(|| T::zero());
let accept_prob = (-(h_new - h_old)).exp();
if rng.random::<f64>() > accept_prob.to_f64().unwrap_or(0.0) {
for k in 0..self.config.input_dim {
q[[j, k]] = q_old[[j, k]];
}
}
}
}
Ok(q)
}
pub fn train_contrastive_divergence(
&mut self,
x: &Array2<T>,
k_steps: usize,
) -> NeuralResult<T> {
let _batch_size = x.nrows();
let energy_pos = self.energy(x)?;
let _grad_pos = self.energy_net.energy_gradient(x)?;
let x_neg = match &self.persistent_chain {
Some(chain)
if matches!(
self.config.training_algorithm,
EBMTrainingAlgorithm::PersistentCD { .. }
) =>
{
let mut chain = chain.clone();
for _ in 0..k_steps {
let grad = self.energy_net.energy_gradient(&chain)?;
let noise = Array2::from_shape_fn(chain.dim(), |_| {
T::from(thread_rng().sample::<f64, _>(
Normal::new(0.0, 1.0).expect("valid distribution params"),
))
.expect("value should be present")
});
chain = &chain - &grad.mapv(|g| g * T::from(0.01).unwrap_or_else(|| T::zero()))
+ &noise.mapv(|n| n * T::from(0.1).unwrap_or_else(|| T::zero()));
}
self.persistent_chain = Some(chain.clone());
chain
}
_ => {
let mut x_neg = x.clone();
for _ in 0..k_steps {
let grad = self.energy_net.energy_gradient(&x_neg)?;
let noise = Array2::from_shape_fn(x_neg.dim(), |_| {
T::from(thread_rng().sample::<f64, _>(
Normal::new(0.0, 1.0).expect("valid distribution params"),
))
.expect("value should be present")
});
x_neg = &x_neg - &grad.mapv(|g| g * T::from(0.01).unwrap_or_else(|| T::zero()))
+ &noise.mapv(|n| n * T::from(0.1).unwrap_or_else(|| T::zero()));
}
x_neg
}
};
let energy_neg = self.energy(&x_neg)?;
let loss = (energy_pos
.mean()
.expect("mean should not fail on non-empty array")
- energy_neg
.mean()
.expect("mean should not fail on non-empty array"))
.abs();
Ok(loss)
}
pub fn config(&self) -> &EBMConfig {
&self.config
}
pub fn num_parameters(&self) -> usize {
self.energy_net.num_parameters()
}
}
#[derive(Debug)]
pub struct HopfieldNetwork<T: FloatBounds> {
weights: Array2<T>,
n_units: usize,
patterns: Vec<Array1<T>>,
}
impl<T: FloatBounds + ScalarOperand> HopfieldNetwork<T> {
pub fn new(n_units: usize) -> Self {
let weights = Array2::zeros((n_units, n_units));
Self {
weights,
n_units,
patterns: Vec::new(),
}
}
pub fn store_pattern(&mut self, pattern: &Array1<T>) -> NeuralResult<()> {
if pattern.len() != self.n_units {
return Err(SklearsError::InvalidParameter {
name: "pattern".to_string(),
reason: format!(
"Pattern length {} does not match network size {}",
pattern.len(),
self.n_units
),
});
}
self.patterns.push(pattern.clone());
for i in 0..self.n_units {
for j in 0..self.n_units {
if i != j {
self.weights[[i, j]] += pattern[i] * pattern[j];
}
}
}
let n_patterns = T::from(self.patterns.len() as f64).unwrap_or_else(|| T::zero());
self.weights.mapv_inplace(|w| w / n_patterns);
Ok(())
}
pub fn recall(&self, initial: &Array1<T>, max_iterations: usize) -> NeuralResult<Array1<T>> {
if initial.len() != self.n_units {
return Err(SklearsError::InvalidParameter {
name: "initial".to_string(),
reason: format!(
"Initial state length {} does not match network size {}",
initial.len(),
self.n_units
),
});
}
let mut state = initial.clone();
for _ in 0..max_iterations {
let mut new_state = state.clone();
for i in 0..self.n_units {
let activation = self.weights.row(i).dot(&state);
new_state[i] = if activation >= T::zero() {
T::one()
} else {
-T::one()
};
}
if new_state == state {
break;
}
state = new_state;
}
Ok(state)
}
pub fn energy(&self, state: &Array1<T>) -> T {
let mut energy = T::zero();
for i in 0..self.n_units {
for j in 0..self.n_units {
energy -= self.weights[[i, j]] * state[i] * state[j];
}
}
energy / T::from(2.0).unwrap_or_else(|| T::zero())
}
pub fn num_patterns(&self) -> usize {
self.patterns.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_energy_network_creation() {
let network: EnergyNetwork<f64> = EnergyNetwork::new(10, vec![32, 16]);
assert_eq!(network.input_dim, 10);
assert!(network.num_parameters() > 0);
}
#[test]
fn test_energy_computation() {
let mut network: EnergyNetwork<f64> = EnergyNetwork::new(8, vec![16]);
let x = Array2::from_shape_fn((4, 8), |(i, j)| (i + j) as f64 * 0.1);
let energy = network.energy(&x).expect("operation should succeed");
assert_eq!(energy.len(), 4);
assert!(energy.iter().all(|&e| e.is_finite()));
}
#[test]
fn test_energy_gradient() {
let mut network: EnergyNetwork<f64> = EnergyNetwork::new(6, vec![12]);
let x = Array2::from_shape_fn((3, 6), |(i, j)| (i + j) as f64 * 0.1);
let grad = network
.energy_gradient(&x)
.expect("operation should succeed");
assert_eq!(grad.dim(), x.dim());
assert!(grad.iter().all(|&g| g.is_finite()));
}
#[test]
fn test_ebm_creation() {
let config = EBMConfig {
input_dim: 10,
hidden_dims: vec![32],
..Default::default()
};
let ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
assert!(ebm.num_parameters() > 0);
}
#[test]
fn test_langevin_sampling() {
let config = EBMConfig {
input_dim: 8,
hidden_dims: vec![16],
sampling_method: SamplingMethod::Langevin {
step_size: 0.01,
num_steps: 10,
},
..Default::default()
};
let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
let samples = ebm.sample(5).expect("sampling should succeed");
assert_eq!(samples.nrows(), 5);
assert_eq!(samples.ncols(), 8);
}
#[test]
fn test_gibbs_sampling() {
let config = EBMConfig {
input_dim: 6,
hidden_dims: vec![12],
sampling_method: SamplingMethod::Gibbs,
..Default::default()
};
let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
let samples = ebm.sample(4).expect("sampling should succeed");
assert_eq!(samples.nrows(), 4);
assert_eq!(samples.ncols(), 6);
assert!(samples.iter().all(|&x| x == 0.0 || x == 1.0));
}
#[test]
fn test_contrastive_divergence() {
let config = EBMConfig {
input_dim: 8,
hidden_dims: vec![16],
..Default::default()
};
let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
let x = Array2::from_shape_fn((4, 8), |(i, j)| (i + j) as f64 * 0.1);
let loss = ebm
.train_contrastive_divergence(&x, 1)
.expect("operation should succeed");
assert!(loss.is_finite());
assert!(loss >= 0.0);
}
#[test]
fn test_hopfield_network_creation() {
let network: HopfieldNetwork<f64> = HopfieldNetwork::new(10);
assert_eq!(network.n_units, 10);
assert_eq!(network.num_patterns(), 0);
}
#[test]
fn test_hopfield_store_pattern() {
let mut network: HopfieldNetwork<f64> = HopfieldNetwork::new(5);
let pattern = Array1::from_vec(vec![1.0, -1.0, 1.0, -1.0, 1.0]);
network
.store_pattern(&pattern)
.expect("operation should succeed");
assert_eq!(network.num_patterns(), 1);
}
#[test]
fn test_hopfield_recall() {
let mut network: HopfieldNetwork<f64> = HopfieldNetwork::new(4);
let pattern = Array1::from_vec(vec![1.0, 1.0, -1.0, -1.0]);
network
.store_pattern(&pattern)
.expect("operation should succeed");
let noisy = Array1::from_vec(vec![1.0, -1.0, -1.0, -1.0]);
let recalled = network
.recall(&noisy, 10)
.expect("operation should succeed");
assert_eq!(recalled.len(), 4);
assert!(recalled.iter().all(|&x| x == 1.0 || x == -1.0));
}
#[test]
fn test_hopfield_energy() {
let mut network: HopfieldNetwork<f64> = HopfieldNetwork::new(4);
let pattern = Array1::from_vec(vec![1.0, -1.0, 1.0, -1.0]);
network
.store_pattern(&pattern)
.expect("operation should succeed");
let energy = network.energy(&pattern);
assert!(energy.is_finite());
let random = Array1::from_vec(vec![1.0, 1.0, 1.0, 1.0]);
let energy_random = network.energy(&random);
assert!(energy < energy_random);
}
#[test]
fn test_hmc_sampling() {
let config = EBMConfig {
input_dim: 6,
hidden_dims: vec![12],
sampling_method: SamplingMethod::HMC {
step_size: 0.01,
num_leapfrog: 5,
},
..Default::default()
};
let mut ebm: EnergyBasedModel<f64> = EnergyBasedModel::new(config);
let samples = ebm.sample(3).expect("sampling should succeed");
assert_eq!(samples.nrows(), 3);
assert_eq!(samples.ncols(), 6);
assert!(samples.iter().all(|&x| x.is_finite()));
}
}