use crate::rng::Pcg;
pub fn gamma_coupling(gamma: f64, t: f64, m: usize) -> f64 {
let e = (-gamma * t).exp();
((1.0 + (m as f64 - 1.0) * e) / (1.0 - e)).ln()
}
pub fn keep_prob(gamma: f64, t: f64, m: usize) -> f64 {
let e = (-gamma * t).exp();
(1.0 + (m as f64 - 1.0) * e) / m as f64
}
pub fn forward_step(x: &mut [i8], gamma: f64, dt: f64, rng: &mut Pcg) {
let keep = keep_prob(gamma, dt, 2);
for s in x.iter_mut() {
if rng.f64() >= keep {
*s = -*s;
}
}
}
pub const G8: [(i64, i64); 2] = [(0, 1), (4, 1)];
pub const G12: [(i64, i64); 3] = [(0, 1), (4, 1), (9, 10)];
pub const G16: [(i64, i64); 4] = [(0, 1), (4, 1), (8, 7), (14, 9)];
pub fn pattern_grid(l: usize, rules: &[(i64, i64)]) -> Vec<(u32, u32)> {
let mut edges = Vec::new();
for y in 0..l as i64 {
for x in 0..l as i64 {
let i = (y * l as i64 + x) as u32;
for &(a, b) in rules {
for (dx, dy) in [(a, b), (-b, a), (-a, -b), (b, -a)] {
let (nx, ny) = (x + dx, y + dy);
if nx >= 0 && ny >= 0 && nx < l as i64 && ny < l as i64 {
let j = (ny * l as i64 + nx) as u32;
if j > i {
edges.push((i, j));
}
}
}
}
}
}
edges.sort_unstable();
edges.dedup();
edges
}
pub struct Ebm {
pub n: usize,
pub edges: Vec<(u16, u16)>,
pub j: Vec<f64>,
pub h: Vec<f64>,
}
impl Ebm {
pub fn new(n: usize, edges: Vec<(u16, u16)>) -> Ebm {
let ne = edges.len();
Ebm { n, edges, j: vec![0.0; ne], h: vec![0.0; n] }
}
pub fn energy(&self, s: &[i8]) -> f64 {
let mut e = 0.0;
for (k, &(a, b)) in self.edges.iter().enumerate() {
e -= self.j[k] * (s[a as usize] * s[b as usize]) as f64;
}
for i in 0..self.n {
e -= self.h[i] * s[i] as f64;
}
e
}
fn field(&self, i: usize, s: &[i8], extra: &[f64]) -> f64 {
let mut f = self.h[i] + extra[i];
for (k, &(a, b)) in self.edges.iter().enumerate() {
if a as usize == i {
f += self.j[k] * s[b as usize] as f64;
} else if b as usize == i {
f += self.j[k] * s[a as usize] as f64;
}
}
f
}
pub fn gibbs(&self, s: &mut [i8], free: &[usize], extra: &[f64], sweeps: usize, rng: &mut Pcg) {
for _ in 0..sweeps {
for &i in free {
let f = self.field(i, s, extra);
let p_up = 1.0 / (1.0 + (-2.0 * f).exp());
s[i] = if rng.f64() < p_up { 1 } else { -1 };
}
}
}
}
pub struct Dtm {
pub steps: Vec<Ebm>,
pub nv: usize,
pub gamma: f64,
pub times: Vec<f64>,
}
impl Dtm {
pub fn new(t_steps: usize, n: usize, nv: usize, edges: Vec<(u16, u16)>, gamma: f64, times: Vec<f64>) -> Dtm {
assert_eq!(times.len(), t_steps + 1);
Dtm {
steps: (0..t_steps).map(|_| Ebm::new(n, edges.clone())).collect(),
nv,
gamma,
times,
}
}
fn clamp_field(&self, t: usize, xt: &[i8]) -> Vec<f64> {
let dt = self.times[t + 1] - self.times[t];
let g = gamma_coupling(self.gamma, dt, 2);
let n = self.steps[0].n;
let mut extra = vec![0.0; n];
for i in 0..self.nv {
extra[i] = 0.5 * g * xt[i] as f64;
}
extra
}
pub fn train_step(
&mut self,
t: usize,
batch: &[(Vec<i8>, Vec<i8>)],
k_sweeps: usize,
lr: f64,
lambda_tc: f64,
rng: &mut Pcg,
) {
let n = self.steps[t].n;
let ne = self.steps[t].edges.len();
let nv = self.nv;
let latents: Vec<usize> = (nv..n).collect();
let all: Vec<usize> = (0..n).collect();
let mut pos_ss = vec![0.0; ne];
let mut pos_s = vec![0.0; n];
let mut neg_ss = vec![0.0; ne];
let mut neg_s = vec![0.0; n];
for (x_prev, x_next) in batch {
let extra = self.clamp_field(t, x_next);
let ebm = &self.steps[t];
let mut s = vec![1i8; n];
s[..nv].copy_from_slice(x_prev);
for i in nv..n {
s[i] = if rng.f64() < 0.5 { 1 } else { -1 };
}
ebm.gibbs(&mut s, &latents, &extra, k_sweeps, rng);
for (k, &(a, b)) in ebm.edges.iter().enumerate() {
pos_ss[k] += (s[a as usize] * s[b as usize]) as f64;
}
for i in 0..n {
pos_s[i] += s[i] as f64;
}
let mut sneg: Vec<i8> = (0..n).map(|_| if rng.f64() < 0.5 { 1 } else { -1 }).collect();
ebm.gibbs(&mut sneg, &all, &extra, k_sweeps, rng);
for (k, &(a, b)) in ebm.edges.iter().enumerate() {
neg_ss[k] += (sneg[a as usize] * sneg[b as usize]) as f64;
}
for i in 0..n {
neg_s[i] += sneg[i] as f64;
}
}
let m = batch.len() as f64;
let ebm = &mut self.steps[t];
for k in 0..ne {
let mut g = -(pos_ss[k] - neg_ss[k]) / m;
if lambda_tc > 0.0 {
let (a, b) = ebm.edges[k];
let (ma, mb) = (neg_s[a as usize] / m, neg_s[b as usize] / m);
g += lambda_tc * -(ma * mb - neg_ss[k] / m);
}
ebm.j[k] -= lr * g;
}
for i in 0..n {
ebm.h[i] -= lr * -(pos_s[i] - neg_s[i]) / m;
}
}
pub fn sample(&self, k_mix: usize, rng: &mut Pcg) -> Vec<i8> {
let n = self.steps[0].n;
let nv = self.nv;
let all: Vec<usize> = (0..n).collect();
let mut x: Vec<i8> = (0..nv).map(|_| if rng.f64() < 0.5 { 1 } else { -1 }).collect();
for t in (0..self.steps.len()).rev() {
let extra = self.clamp_field(t, &x);
let mut s: Vec<i8> = (0..n).map(|_| if rng.f64() < 0.5 { 1 } else { -1 }).collect();
self.steps[t].gibbs(&mut s, &all, &extra, k_mix, rng);
x.copy_from_slice(&s[..nv]);
}
x
}
pub fn exact_log_cond(&self, t: usize, x_prev: &[i8], x_next: &[i8]) -> f64 {
let n = self.steps[t].n;
let nv = self.nv;
let nl = n - nv;
let extra = self.clamp_field(t, x_next);
let ebm = &self.steps[t];
let full_e = |xv: &[i8], zm: usize| -> f64 {
let mut s = vec![0i8; n];
s[..nv].copy_from_slice(xv);
for l in 0..nl {
s[nv + l] = if zm >> l & 1 == 1 { 1 } else { -1 };
}
let mut e = ebm.energy(&s);
for i in 0..n {
e -= extra[i] * s[i] as f64;
}
e
};
let mut num = 0.0f64;
for zm in 0..(1usize << nl) {
num += (-full_e(x_prev, zm)).exp();
}
let mut den = 0.0f64;
let mut xv = vec![0i8; nv];
for xm in 0..(1usize << nv) {
for b in 0..nv {
xv[b] = if xm >> b & 1 == 1 { 1 } else { -1 };
}
for zm in 0..(1usize << nl) {
den += (-full_e(&xv, zm)).exp();
}
}
num.ln() - den.ln()
}
pub fn exact_nll(&self, data: &[Vec<i8>]) -> f64 {
let nv = self.nv;
let t_steps = self.steps.len();
let mut nll = 0.0;
let masks = 1usize << nv;
let to_x = |m: usize| -> Vec<i8> {
(0..nv).map(|b| if m >> b & 1 == 1 { 1 } else { -1 }).collect()
};
let step_q = |xa: &[i8], xb: &[i8], dt: f64| -> f64 {
let keep = keep_prob(self.gamma, dt, 2);
let mut q = 1.0;
for i in 0..nv {
q *= if xa[i] == xb[i] { keep } else { 1.0 - keep };
}
q
};
let mut traj = vec![0usize; t_steps + 1];
loop {
let x0 = to_x(traj[0]);
let mut q = data
.iter()
.map(|d| if d[..] == x0[..] { 1.0 / data.len() as f64 } else { 0.0 })
.sum::<f64>();
if q > 0.0 {
for t in 0..t_steps {
let xa = to_x(traj[t]);
let xb = to_x(traj[t + 1]);
q *= step_q(&xa, &xb, self.times[t + 1] - self.times[t]);
}
if q > 0.0 {
let mut lp = 0.0;
for t in 0..t_steps {
lp += self.exact_log_cond(t, &to_x(traj[t]), &to_x(traj[t + 1]));
}
nll -= q * lp;
}
}
let mut c = 0;
loop {
traj[c] += 1;
if traj[c] < masks {
break;
}
traj[c] = 0;
c += 1;
if c > t_steps {
return nll;
}
}
}
}
}
pub fn acp_update(
lambda: f64,
a_m: f64,
a_prev: Option<f64>,
eps: f64,
delta: f64,
lambda_min: f64,
) -> f64 {
let lp = lambda.max(lambda_min);
let out = if a_m < eps {
(1.0 - delta) * lp
} else if a_prev.is_none() || a_m <= a_prev.unwrap() {
lp
} else {
(1.0 + delta) * lp
};
if out < lambda_min {
0.0
} else {
out
}
}
pub fn autocorr(series: &[f64], k: usize) -> f64 {
let n = series.len();
assert!(k < n);
let mu = series.iter().sum::<f64>() / n as f64;
let var = series.iter().map(|v| (v - mu) * (v - mu)).sum::<f64>() / n as f64;
if var == 0.0 {
return 0.0;
}
let mut c = 0.0;
for i in 0..n - k {
c += (series[i] - mu) * (series[i + k] - mu);
}
c / ((n - k) as f64 * var)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn forward_sign_trap() {
for gamma in [0.3, 1.0, 2.5] {
for t in [0.1, 0.5, 1.0, 3.0] {
let g = gamma_coupling(gamma, t, 2);
let keep_closed = keep_prob(gamma, t, 2);
let keep_energy = 1.0 / (1.0 + (-g).exp());
assert!(
(keep_energy - keep_closed).abs() < 1e-14,
"gamma {gamma} t {t}: energy-form keep {keep_energy} vs closed {keep_closed}"
);
let wrong = 1.0 / (1.0 + g.exp());
assert!((wrong - (1.0 - keep_closed)).abs() < 1e-14);
if keep_closed - 0.5 > 1e-3 {
assert!((wrong - keep_closed).abs() > 1e-3, "trap not discriminable at gamma {gamma}, t {t}");
}
}
}
}
#[test]
fn kernel_semigroup() {
for gamma in [0.4, 1.3] {
for (t1, t2) in [(0.2, 0.7), (0.5, 0.5), (1.0, 2.0)] {
let k1 = keep_prob(gamma, t1, 2);
let k2 = keep_prob(gamma, t2, 2);
let k12 = keep_prob(gamma, t1 + t2, 2);
let composed = k1 * k2 + (1.0 - k1) * (1.0 - k2);
assert!((composed - k12).abs() < 1e-14);
}
}
}
#[test]
fn pattern_grid_structure() {
for (rules, deg) in [(&G8[..], 8usize), (&G12[..], 12), (&G16[..], 16)] {
let l = 40usize;
let edges = pattern_grid(l, rules);
let mut count = vec![0usize; l * l];
for &(a, b) in &edges {
count[a as usize] += 1;
count[b as usize] += 1;
let (ax, ay) = (a as usize % l, a as usize / l);
let (bx, by) = (b as usize % l, b as usize / l);
assert_eq!((ax + ay + bx + by) % 2, 1, "edge does not flip parity");
}
let i = (l / 2) * l + l / 2;
assert_eq!(count[i], deg, "interior degree for rule set of {} rules", rules.len());
}
}
#[test]
fn acp_scripted_sequence() {
let (eps, delta, lmin) = (0.03, 0.2, 1e-4);
let a = [0.5, 0.6, 0.4, 0.01];
let mut lambda = 0.01;
let mut prev: Option<f64> = None;
let mut traj = Vec::new();
for (m, &am) in a.iter().enumerate() {
let ap = if m == 0 { None } else { prev };
lambda = acp_update(lambda, am, ap, eps, delta, lmin);
traj.push(lambda);
prev = Some(am);
}
let want = [0.01, 0.012, 0.012, 0.0096];
for (got, want) in traj.iter().zip(&want) {
assert!((got - want).abs() < 1e-12, "traj {:?} vs {:?}", traj, want);
}
}
#[test]
fn autocorr_alternating_exact() {
let series: Vec<f64> = (0..1000).map(|i| if i % 2 == 0 { 1.0 } else { -1.0 }).collect();
for k in 0..5 {
let want = if k % 2 == 0 { 1.0 } else { -1.0 };
assert!((autocorr(&series, k) - want).abs() < 1e-12);
}
}
#[test]
fn exact_enumeration_gradient_check() {
let edges: Vec<(u16, u16)> = vec![(0, 1), (1, 2), (0, 3), (1, 3), (2, 4), (1, 4), (3, 4)];
let mut dtm = Dtm::new(2, 5, 3, edges.clone(), 1.0, vec![0.0, 1.0, 2.0]);
let mut rng = Pcg::new(0x601D, 1);
for t in 0..2 {
for j in dtm.steps[t].j.iter_mut() {
*j = (rng.f64() - 0.5) * 0.6;
}
for h in dtm.steps[t].h.iter_mut() {
*h = (rng.f64() - 0.5) * 0.4;
}
}
let data = vec![vec![1i8, 1, -1], vec![-1, 1, 1]];
let nv = 3usize;
let nl = 2usize;
let to_x = |m: usize| -> Vec<i8> {
(0..nv).map(|b| if m >> b & 1 == 1 { 1 } else { -1 }).collect()
};
let keep = |dt: f64| keep_prob(1.0, dt, 2);
let q_step = |xa: &[i8], xb: &[i8], dt: f64| -> f64 {
let k = keep(dt);
(0..nv).map(|i| if xa[i] == xb[i] { k } else { 1.0 - k }).product()
};
let q_pair = |t: usize, a: usize, b: usize| -> f64 {
let xa = to_x(a);
let xb = to_x(b);
match t {
0 => {
let q0: f64 = data
.iter()
.map(|d| if d[..] == xa[..] { 0.5 } else { 0.0 })
.sum();
q0 * q_step(&xa, &xb, 1.0)
}
_ => {
let mut q1 = 0.0;
for m0 in 0..8usize {
let x0 = to_x(m0);
let q0: f64 = data
.iter()
.map(|d| if d[..] == x0[..] { 0.5 } else { 0.0 })
.sum();
q1 += q0 * q_step(&x0, &xa, 1.0);
}
q1 * q_step(&xa, &xb, 1.0)
}
}
};
let stat = |dtm: &Dtm, t: usize, clamp_prev: Option<usize>, b_mask: usize| -> (Vec<f64>, Vec<f64>) {
let ebm = &dtm.steps[t];
let extra = dtm.clamp_field(t, &to_x(b_mask));
let mut zsum = 0.0;
let mut ess = vec![0.0; ebm.edges.len()];
let mut es = vec![0.0; ebm.n];
let xs: Vec<usize> = match clamp_prev {
Some(a) => vec![a],
None => (0..8).collect(),
};
for &xm in &xs {
let xv = to_x(xm);
for zm in 0..(1usize << nl) {
let mut s = vec![0i8; 5];
s[..nv].copy_from_slice(&xv);
for l in 0..nl {
s[nv + l] = if zm >> l & 1 == 1 { 1 } else { -1 };
}
let mut e = ebm.energy(&s);
for i in 0..5 {
e -= extra[i] * s[i] as f64;
}
let w = (-e).exp();
zsum += w;
for (k, &(a, b)) in ebm.edges.iter().enumerate() {
ess[k] += w * (s[a as usize] * s[b as usize]) as f64;
}
for i in 0..5 {
es[i] += w * s[i] as f64;
}
}
}
for v in ess.iter_mut() {
*v /= zsum;
}
for v in es.iter_mut() {
*v /= zsum;
}
(ess, es)
};
let ne = edges.len();
for t in 0..2usize {
let mut g_j = vec![0.0; ne];
let mut g_h = vec![0.0; 5];
for b_mask in 0..8usize {
let mut qb = 0.0;
for a_mask in 0..8usize {
let q = q_pair(t, a_mask, b_mask);
if q == 0.0 {
continue;
}
qb += q;
let (ess, es) = stat(&dtm, t, Some(a_mask), b_mask);
for k in 0..ne {
g_j[k] += q * -ess[k];
}
for i in 0..5 {
g_h[i] += q * -es[i];
}
}
if qb == 0.0 {
continue;
}
let (ess, es) = stat(&dtm, t, None, b_mask);
for k in 0..ne {
g_j[k] -= qb * -ess[k];
}
for i in 0..5 {
g_h[i] -= qb * -es[i];
}
}
let fd = 1e-6;
for k in 0..ne {
let orig = dtm.steps[t].j[k];
dtm.steps[t].j[k] = orig + fd;
let up = dtm.exact_nll(&data);
dtm.steps[t].j[k] = orig - fd;
let dn = dtm.exact_nll(&data);
dtm.steps[t].j[k] = orig;
let want = (up - dn) / (2.0 * fd);
assert!(
(g_j[k] - want).abs() < 1e-7,
"t {t} J[{k}]: eq14 {} vs FD {}",
g_j[k],
want
);
}
for i in 0..5 {
let orig = dtm.steps[t].h[i];
dtm.steps[t].h[i] = orig + fd;
let up = dtm.exact_nll(&data);
dtm.steps[t].h[i] = orig - fd;
let dn = dtm.exact_nll(&data);
dtm.steps[t].h[i] = orig;
let want = (up - dn) / (2.0 * fd);
assert!(
(g_h[i] - want).abs() < 1e-7,
"t {t} h[{i}]: eq14 {} vs FD {}",
g_h[i],
want
);
}
}
}
#[test]
fn training_reduces_exact_nll() {
let edges: Vec<(u16, u16)> = vec![(0, 1), (1, 2), (2, 3), (0, 4), (1, 4), (2, 5), (3, 5)];
let mut dtm = Dtm::new(2, 6, 4, edges, 1.2, vec![0.0, 0.8, 1.6]);
let data = vec![vec![1i8, 1, 1, 1], vec![-1, -1, -1, -1]];
let nll0 = dtm.exact_nll(&data);
let mut rng = Pcg::new(0x7247, 3);
for _iter in 0..150 {
for t in 0..2 {
let mut batch = Vec::new();
for _ in 0..24 {
let d = &data[(rng.f64() * 2.0) as usize % 2];
let mut xa = d.clone();
if t > 0 {
forward_step(&mut xa, dtm.gamma, dtm.times[t] - dtm.times[0], &mut rng);
}
let mut xb = xa.clone();
forward_step(&mut xb, dtm.gamma, dtm.times[t + 1] - dtm.times[t], &mut rng);
batch.push((xa, xb));
}
dtm.train_step(t, &batch, 25, 0.05, 0.0, &mut rng);
}
}
let nll1 = dtm.exact_nll(&data);
assert!(
nll1 < nll0 - 0.5,
"training did not reduce exact NLL: {nll0:.3} -> {nll1:.3}"
);
}
}