use crate::gibbs::Sampler;
use crate::graph::{Graph, GraphBuilder};
use crate::rng::Pcg;
#[derive(Clone, Debug)]
pub struct Dataset {
pub visible: usize,
pub rows: Vec<Vec<i8>>,
}
#[derive(Clone, Debug, PartialEq)]
pub enum Error {
NoData,
RowWidth { row: usize, len: usize, want: usize },
NotASpin { row: usize, at: usize, value: i8 },
TooSmall { spins: usize, visible: usize },
}
impl core::fmt::Display for Error {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Error::NoData => write!(f, "no data rows; there is nothing to fit"),
Error::RowWidth { row, len, want } => {
write!(f, "row {row} has {len} visible entries, and the dataset declares {want}")
}
Error::NotASpin { row, at, value } => {
write!(f, "row {row} position {at} is {value}, and a spin is -1 or +1")
}
Error::TooSmall { spins, visible } => {
write!(f, "the model has {spins} spins and the data needs {visible} visible")
}
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Params {
pub epochs: usize,
pub k: usize,
pub positive_sweeps: usize,
pub learning_rate: f64,
pub batch: usize,
}
impl Default for Params {
fn default() -> Self {
Params { epochs: 300, k: 5, positive_sweeps: 5, learning_rate: 0.05, batch: 8 }
}
}
pub struct Trained {
pub graph: Graph,
pub log_likelihood: Option<f64>,
pub epochs_run: usize,
}
impl core::fmt::Debug for Trained {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Trained")
.field("spins", &self.graph.n)
.field("edges", &self.graph.n_edges)
.field("log_likelihood", &self.log_likelihood)
.field("epochs_run", &self.epochs_run)
.finish()
}
}
pub fn train(structure: &Graph, data: &Dataset, p: &Params, seed: u64) -> Result<Trained, Error> {
check(structure, data)?;
let n = structure.n;
let mut rng = Pcg::new(seed, 0x00EB_3600);
let mut edges: Vec<(usize, usize, f64)> = Vec::with_capacity(structure.n_edges);
for i in 0..n {
for k in structure.offset[i]..structure.offset[i + 1] {
let j = structure.nbr[k] as usize;
if j > i {
edges.push((i, j, structure.w[k]));
}
}
}
let mut bias: Vec<f64> = structure.h.clone();
let build = |edges: &[(usize, usize, f64)], bias: &[f64]| {
let mut gb = GraphBuilder::new(n);
for &(i, j, w) in edges {
gb.couple(i, j, w);
}
for (i, &b) in bias.iter().enumerate() {
if b != 0.0 {
gb.bias(i, b);
}
}
gb.build()
};
let mut order: Vec<usize> = (0..data.rows.len()).collect();
for epoch in 0..p.epochs {
for i in (1..order.len()).rev() {
let j = (rng.f64() * (i + 1) as f64) as usize % (i + 1);
order.swap(i, j);
}
let g = build(&edges, &bias);
let decay = if p.epochs > 1 {
1.0 - 0.9 * epoch as f64 / (p.epochs - 1) as f64
} else {
1.0
};
for chunk in order.chunks(p.batch.max(1)) {
let mut d_edge = vec![0.0f64; edges.len()];
let mut d_bias = vec![0.0f64; n];
for &r in chunk {
let row = &data.rows[r];
let seed = (rng.next_u32() as u64) << 32 | rng.next_u32() as u64;
let mut smp = Sampler::new(&g, 1.0, seed);
for (i, &v) in row.iter().enumerate() {
smp.clamp(i, v);
}
smp.sweeps(p.positive_sweeps.max(1), None);
let pos = smp.s.clone();
for i in 0..data.visible {
smp.unclamp(i);
}
smp.sweeps(p.k.max(1), None);
let neg = &smp.s;
for (e, &(i, j, _)) in edges.iter().enumerate() {
d_edge[e] += (pos[i] * pos[j]) as f64 - (neg[i] * neg[j]) as f64;
}
for i in 0..n {
d_bias[i] += pos[i] as f64 - neg[i] as f64;
}
}
let scale = p.learning_rate * decay / chunk.len() as f64;
for (e, w) in edges.iter_mut().enumerate() {
w.2 += scale * d_edge[e];
}
for i in 0..n {
bias[i] += scale * d_bias[i];
}
}
}
let graph = build(&edges, &bias);
let log_likelihood = exact_log_likelihood(&graph, data).ok();
Ok(Trained { graph, log_likelihood, epochs_run: p.epochs })
}
pub const MAX_ENUMERATED: usize = 22;
pub fn exact_log_likelihood(g: &Graph, data: &Dataset) -> Result<f64, Error> {
check(g, data)?;
if g.n > MAX_ENUMERATED {
return Err(Error::TooSmall { spins: g.n, visible: data.visible });
}
let mut max_neg_e = f64::NEG_INFINITY;
let states = 1usize << g.n;
let mut energies = Vec::with_capacity(states);
let mut s = vec![-1i8; g.n];
for mask in 0..states {
for i in 0..g.n {
s[i] = if mask >> i & 1 == 1 { 1 } else { -1 };
}
let e = -g.energy(&s);
max_neg_e = max_neg_e.max(e);
energies.push(e);
}
let z: f64 = energies.iter().map(|e| (e - max_neg_e).exp()).sum();
let log_z = max_neg_e + z.ln();
let vmask = (1usize << data.visible) - 1;
let mut per_visible = vec![0.0f64; 1usize << data.visible];
for (mask, &e) in energies.iter().enumerate() {
per_visible[mask & vmask] += (e - max_neg_e).exp();
}
let mut total = 0.0;
for row in &data.rows {
let mut key = 0usize;
for (i, &v) in row.iter().enumerate() {
if v == 1 {
key |= 1 << i;
}
}
total += max_neg_e + per_visible[key].ln() - log_z;
}
Ok(total / data.rows.len() as f64)
}
fn check(g: &Graph, data: &Dataset) -> Result<(), Error> {
if data.rows.is_empty() {
return Err(Error::NoData);
}
if g.n < data.visible {
return Err(Error::TooSmall { spins: g.n, visible: data.visible });
}
for (r, row) in data.rows.iter().enumerate() {
if row.len() != data.visible {
return Err(Error::RowWidth { row: r, len: row.len(), want: data.visible });
}
if let Some(at) = row.iter().position(|&v| v != 1 && v != -1) {
return Err(Error::NotASpin { row: r, at, value: row[at] });
}
}
Ok(())
}
pub fn rbm(visible: usize, hidden: usize) -> Graph {
let mut gb = GraphBuilder::new(visible + hidden);
for v in 0..visible {
for h in 0..hidden {
gb.couple(v, visible + h, 0.0);
}
}
gb.build()
}
pub fn dbm(visible: usize, hidden: &[usize]) -> Graph {
let n = visible + hidden.iter().sum::<usize>();
let mut gb = GraphBuilder::new(n);
let mut below = (0..visible).collect::<Vec<_>>();
let mut next = visible;
for &w in hidden {
let layer: Vec<usize> = (next..next + w).collect();
for &a in &below {
for &b in &layer {
gb.couple(a, b, 0.0);
}
}
next += w;
below = layer;
}
gb.build()
}
pub fn bars_and_stripes(side: usize) -> Dataset {
let n = side * side;
let mut seen: Vec<Vec<i8>> = Vec::new();
for mask in 0..(1usize << side) {
for stripes in [false, true] {
let mut row = vec![-1i8; n];
for a in 0..side {
if mask >> a & 1 == 1 {
for b in 0..side {
row[if stripes { a * side + b } else { b * side + a }] = 1;
}
}
}
if !seen.contains(&row) {
seen.push(row);
}
}
}
Dataset { visible: n, rows: seen }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_fully_visible_fit_matches_the_data_correlations() {
let mut gb = GraphBuilder::new(3);
gb.couple(0, 1, 0.0);
gb.couple(0, 2, 0.0);
gb.couple(1, 2, 0.0);
let structure = gb.build();
let rows: Vec<Vec<i8>> = vec![
vec![1, 1, 1],
vec![1, 1, 1],
vec![1, 1, -1],
vec![-1, -1, 1],
vec![-1, -1, 1],
vec![-1, -1, -1],
];
let data = Dataset { visible: 3, rows: rows.clone() };
let p = Params { epochs: 4_000, k: 20, learning_rate: 0.05, batch: 6, positive_sweeps: 1 };
let t = train(&structure, &data, &p, 7).unwrap();
let m = rows.len() as f64;
let dc = |i: usize, j: usize| {
rows.iter().map(|r| (r[i] * r[j]) as f64).sum::<f64>() / m
};
let dm = |i: usize| rows.iter().map(|r| r[i] as f64).sum::<f64>() / m;
let g = &t.graph;
let mut z = 0.0;
let mut corr = [[0.0f64; 3]; 3];
let mut mag = [0.0f64; 3];
for mask in 0..8usize {
let s: Vec<i8> = (0..3).map(|i| if mask >> i & 1 == 1 { 1 } else { -1 }).collect();
let w = (-g.energy(&s)).exp();
z += w;
for i in 0..3 {
mag[i] += w * s[i] as f64;
for j in 0..3 {
corr[i][j] += w * (s[i] * s[j]) as f64;
}
}
}
for i in 0..3 {
assert!(
(mag[i] / z - dm(i)).abs() < 0.05,
"magnetisation {i}: model {:.4} vs data {:.4}",
mag[i] / z,
dm(i)
);
for j in (i + 1)..3 {
assert!(
(corr[i][j] / z - dc(i, j)).abs() < 0.05,
"correlation ({i},{j}): model {:.4} vs data {:.4}",
corr[i][j] / z,
dc(i, j)
);
}
}
}
#[test]
fn training_raises_the_exact_log_likelihood() {
let data = bars_and_stripes(2);
let structure = rbm(4, 4);
let before = exact_log_likelihood(&structure, &data).unwrap();
assert!((before - (-4.0 * 2f64.ln())).abs() < 1e-9, "{before}");
let p = Params { epochs: 600, k: 10, ..Params::default() };
let t = train(&structure, &data, &p, 3).unwrap();
let after = t.log_likelihood.unwrap();
assert!(after > before + 0.05, "training must help: {before:.4} -> {after:.4}");
assert!(after < 0.0, "a log-likelihood is negative: {after}");
}
#[test]
fn bars_and_stripes_is_the_right_set() {
for side in [2usize, 3, 4] {
let d = bars_and_stripes(side);
assert_eq!(d.visible, side * side);
assert_eq!(d.rows.len(), 2 * (1 << side) - 2, "side {side}");
assert!(d.rows.iter().all(|r| r.iter().all(|&v| v == 1 || v == -1)));
}
let d = bars_and_stripes(3);
for r in &d.rows {
let rows_const = (0..3).all(|a| (0..3).all(|b| r[a * 3 + b] == r[a * 3]));
let cols_const = (0..3).all(|a| (0..3).all(|b| r[b * 3 + a] == r[a]));
assert!(rows_const || cols_const, "{r:?}");
}
}
#[test]
fn a_deep_machine_has_fewer_edges_than_a_wide_one_with_the_same_latents() {
let wide = rbm(9, 8);
let deep = dbm(9, &[4, 4]);
assert_eq!(wide.n, deep.n);
assert_eq!(wide.n_edges, 9 * 8);
assert_eq!(deep.n_edges, 9 * 4 + 4 * 4);
assert!(deep.n_edges < wide.n_edges);
assert_eq!(dbm(9, &[8]).n_edges, wide.n_edges);
}
#[test]
fn a_malformed_dataset_is_refused_by_name() {
let g = rbm(3, 2);
let p = Params::default();
let err = |d: Dataset| train(&g, &d, &p, 1).unwrap_err();
assert_eq!(err(Dataset { visible: 3, rows: vec![] }), Error::NoData);
assert_eq!(
err(Dataset { visible: 3, rows: vec![vec![1, 1]] }),
Error::RowWidth { row: 0, len: 2, want: 3 }
);
assert_eq!(
err(Dataset { visible: 3, rows: vec![vec![1, 0, 1]] }),
Error::NotASpin { row: 0, at: 1, value: 0 }
);
assert_eq!(
err(Dataset { visible: 9, rows: vec![vec![1; 9]] }),
Error::TooSmall { spins: 5, visible: 9 }
);
}
#[test]
fn an_enumeration_too_large_is_refused_rather_than_attempted() {
let g = rbm(20, 10);
let d = Dataset { visible: 20, rows: vec![vec![1i8; 20]] };
assert!(matches!(exact_log_likelihood(&g, &d), Err(Error::TooSmall { .. })));
}
}