use crate::graph::GraphBuilder;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Encoding {
OneHot,
Binary,
DomainWall,
}
impl Encoding {
pub fn spins(&self, k: usize) -> usize {
assert!(k >= 2, "a variable with fewer than 2 values is a constant");
match self {
Encoding::OneHot => k,
Encoding::Binary => (usize::BITS - (k - 1).leading_zeros()) as usize,
Encoding::DomainWall => k - 1,
}
}
pub fn penalty_couplings(&self, k: usize) -> usize {
match self {
Encoding::OneHot => k * (k - 1) / 2,
Encoding::Binary => 0,
Encoding::DomainWall => k.saturating_sub(2),
}
}
pub fn is_exact(&self, k: usize) -> bool {
match self {
Encoding::Binary => k.is_power_of_two(),
_ => true,
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct Slot {
pub base: usize,
pub k: usize,
pub encoding: Encoding,
}
impl Slot {
pub fn new(base: usize, k: usize, encoding: Encoding) -> Self {
assert!(k >= 2, "a variable with fewer than 2 values is a constant");
Slot { base, k, encoding }
}
pub fn width(&self) -> usize {
self.encoding.spins(self.k)
}
pub fn range(&self) -> std::ops::Range<usize> {
self.base..self.base + self.width()
}
pub fn encode(&self, value: usize, s: &mut [i8]) {
assert!(value < self.k, "value {value} out of range for a {}-valued variable", self.k);
let w = self.width();
match self.encoding {
Encoding::OneHot => {
for i in 0..w {
s[self.base + i] = if i == value { 1 } else { -1 };
}
}
Encoding::Binary => {
for i in 0..w {
s[self.base + i] = if (value >> i) & 1 == 1 { 1 } else { -1 };
}
}
Encoding::DomainWall => {
for i in 0..w {
s[self.base + i] = if i < value { 1 } else { -1 };
}
}
}
}
pub fn decode(&self, s: &[i8]) -> Option<usize> {
let w = self.width();
let bits = &s[self.base..self.base + w];
match self.encoding {
Encoding::OneHot => {
let mut found = None;
for (i, &b) in bits.iter().enumerate() {
if b > 0 {
if found.is_some() {
return None; }
found = Some(i);
}
}
found
}
Encoding::Binary => {
let mut v = 0usize;
for (i, &b) in bits.iter().enumerate() {
if b > 0 {
v |= 1 << i;
}
}
if v < self.k {
Some(v)
} else {
None }
}
Encoding::DomainWall => {
let v = bits.iter().take_while(|&&b| b > 0).count();
if bits[v..].iter().all(|&b| b < 0) {
Some(v)
} else {
None }
}
}
}
pub fn add_penalty(&self, b: &mut GraphBuilder, p: f64) -> bool {
let w = self.width();
match self.encoding {
Encoding::OneHot => {
for i in 0..w {
for j in (i + 1)..w {
b.couple(self.base + i, self.base + j, -p / 2.0);
}
b.bias(self.base + i, -p * (self.k as f64 - 2.0) / 2.0);
}
true
}
Encoding::Binary => self.k.is_power_of_two(),
Encoding::DomainWall => {
for i in 0..w.saturating_sub(1) {
b.couple(self.base + i, self.base + i + 1, p);
}
b.bias(self.base, p);
b.bias(self.base + w - 1, -p);
true
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::GraphBuilder;
const ALL: [Encoding; 3] = [Encoding::OneHot, Encoding::Binary, Encoding::DomainWall];
#[test]
fn encode_decode_round_trips() {
for enc in ALL {
for k in 2..=9 {
let slot = Slot::new(0, k, enc);
let mut s = vec![0i8; slot.width()];
for v in 0..k {
slot.encode(v, &mut s);
assert_eq!(slot.decode(&s), Some(v), "{enc:?} k={k} v={v}");
}
}
}
}
#[test]
fn domain_wall_uses_fewer_spins_than_one_hot() {
for k in 2..=32 {
assert!(
Encoding::DomainWall.spins(k) < Encoding::OneHot.spins(k),
"k={k}"
);
}
}
#[test]
fn domain_wall_penalty_is_linear_where_one_hot_is_quadratic() {
for k in [4, 8, 16, 64] {
let dw = Encoding::DomainWall.penalty_couplings(k);
let oh = Encoding::OneHot.penalty_couplings(k);
assert_eq!(dw, k - 2);
assert_eq!(oh, k * (k - 1) / 2);
assert!(dw < oh, "k={k}: dw {dw} vs oh {oh}");
}
assert_eq!(Encoding::DomainWall.penalty_couplings(1000), 998);
}
fn ground_states(enc: Encoding, k: usize, p: f64) -> Vec<usize> {
let slot = Slot::new(0, k, enc);
let w = slot.width();
let mut b = GraphBuilder::new(w);
slot.add_penalty(&mut b, p);
let g = b.build();
let mut best = f64::INFINITY;
let mut at_best = Vec::new();
for mask in 0..(1usize << w) {
let s: Vec<i8> = (0..w).map(|i| if mask >> i & 1 == 1 { 1 } else { -1 }).collect();
let e = g.energy(&s);
if e < best - 1e-9 {
best = e;
at_best.clear();
}
if e < best + 1e-9 {
at_best.push(mask);
}
}
at_best
}
#[test]
fn domain_wall_ground_states_are_exactly_the_codewords() {
for k in 2..=8 {
let slot = Slot::new(0, k, Encoding::DomainWall);
let g = ground_states(Encoding::DomainWall, k, 2.0);
assert_eq!(g.len(), k, "k={k}: expected {k} degenerate ground states, got {}", g.len());
let w = slot.width();
let mut seen: Vec<usize> = g
.iter()
.map(|&mask| {
let s: Vec<i8> =
(0..w).map(|i| if mask >> i & 1 == 1 { 1 } else { -1 }).collect();
slot.decode(&s).expect("a ground state must be a valid codeword")
})
.collect();
seen.sort_unstable();
assert_eq!(seen, (0..k).collect::<Vec<_>>(), "k={k}");
}
}
#[test]
fn one_hot_ground_states_are_exactly_the_codewords() {
for k in 2..=7 {
let g = ground_states(Encoding::OneHot, k, 2.0);
assert_eq!(g.len(), k, "k={k}");
}
}
#[test]
fn binary_is_honest_about_surplus_codes() {
let slot = Slot::new(0, 6, Encoding::Binary);
assert_eq!(slot.width(), 3);
assert!(!Encoding::Binary.is_exact(6));
let mut b = GraphBuilder::new(3);
assert!(!slot.add_penalty(&mut b, 1.0), "must report that it cannot be exact");
let mut s = vec![-1i8; 3];
s[1] = 1;
s[2] = 1; assert_eq!(slot.decode(&s), None, "a surplus code must decode to None, not a guess");
assert!(Encoding::Binary.is_exact(8));
let mut b = GraphBuilder::new(3);
assert!(Slot::new(0, 8, Encoding::Binary).add_penalty(&mut b, 1.0));
}
#[test]
fn invalid_states_decode_to_none_rather_than_a_guess() {
let oh = Slot::new(0, 4, Encoding::OneHot);
assert_eq!(oh.decode(&[1, 1, -1, -1]), None);
assert_eq!(oh.decode(&[-1, -1, -1, -1]), None, "none hot is also invalid");
let dw = Slot::new(0, 5, Encoding::DomainWall);
assert_eq!(dw.decode(&[1, -1, 1, -1]), None);
assert_eq!(dw.decode(&[1, 1, -1, -1]), Some(2), "one wall is fine");
}
#[test]
fn slots_can_be_packed_side_by_side() {
let a = Slot::new(0, 4, Encoding::DomainWall); let b = Slot::new(a.range().end, 5, Encoding::OneHot); assert_eq!(a.range(), 0..3);
assert_eq!(b.range(), 3..8);
let mut s = vec![0i8; 8];
a.encode(2, &mut s);
b.encode(3, &mut s);
assert_eq!(a.decode(&s), Some(2));
assert_eq!(b.decode(&s), Some(3));
}
}