use serde::{Deserialize, Serialize};
use crate::error::TopologyError;
use crate::record::{RecordSet, DEFAULT_SURVIVOR_THRESHOLD};
use crate::topology::Topology;
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CodebookConfig {
pub capacity: usize,
pub survivor_threshold: f32,
pub max_iterations: usize,
}
impl Default for CodebookConfig {
fn default() -> Self {
Self {
capacity: 16,
survivor_threshold: DEFAULT_SURVIVOR_THRESHOLD,
max_iterations: 32,
}
}
}
impl CodebookConfig {
fn validate(&self) -> Result<(), TopologyError> {
if self.capacity == 0 {
return Err(TopologyError::BadConfig {
field: "capacity",
expected: "at least 1",
found: "0".into(),
});
}
if !self.survivor_threshold.is_finite() {
return Err(TopologyError::BadConfig {
field: "survivor_threshold",
expected: "finite",
found: format!("{}", self.survivor_threshold),
});
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(try_from = "CodebookWire")]
pub struct Codebook {
n: usize,
codes: Vec<Topology>,
occupancy: Vec<usize>,
capacity: usize,
distinct_survivors: usize,
}
#[derive(Deserialize)]
struct CodebookWire {
n: usize,
codes: Vec<Topology>,
occupancy: Vec<usize>,
capacity: usize,
distinct_survivors: usize,
}
impl TryFrom<CodebookWire> for Codebook {
type Error = TopologyError;
fn try_from(wire: CodebookWire) -> Result<Self, Self::Error> {
let book = Codebook {
n: wire.n,
codes: wire.codes,
occupancy: wire.occupancy,
capacity: wire.capacity,
distinct_survivors: wire.distinct_survivors,
};
book.validate()?;
Ok(book)
}
}
impl Codebook {
pub fn fit(records: &RecordSet, config: &CodebookConfig) -> Result<Self, TopologyError> {
config.validate()?;
let survivor_indices = records.survivors(config.survivor_threshold);
if survivor_indices.is_empty() {
return Err(TopologyError::NoRecords { kind: "surviving" });
}
let mut distinct: Vec<Topology> = Vec::new();
for &i in &survivor_indices {
let t = &records.records()[i].topology;
if !distinct.iter().any(|d| d == t) {
distinct.push(t.clone());
}
}
distinct.sort_by(|a, b| a.key().cmp(&b.key()));
let distinct_survivors = distinct.len();
let n = records.team_size();
if distinct.len() <= config.capacity {
let occupancy = vec![1; distinct.len()];
return Ok(Self {
n,
codes: distinct,
occupancy,
capacity: config.capacity,
distinct_survivors,
});
}
let codes = k_medoids(&distinct, config.capacity, config.max_iterations)?;
let mut occupancy = vec![0usize; codes.len()];
for t in &distinct {
occupancy[nearest(&codes, t)?] += 1;
}
Ok(Self {
n,
codes,
occupancy,
capacity: config.capacity,
distinct_survivors,
})
}
pub fn from_topologies(topologies: Vec<Topology>) -> Result<Self, TopologyError> {
if topologies.is_empty() {
return Err(TopologyError::EmptyCodebook);
}
let n = topologies[0].n();
let mut codes: Vec<Topology> = Vec::new();
for t in topologies {
if t.n() != n {
return Err(TopologyError::SizeMismatch {
expected: n,
found: t.n(),
});
}
if !codes.iter().any(|c| c == &t) {
codes.push(t);
}
}
let occupancy = vec![1; codes.len()];
let capacity = codes.len();
let distinct_survivors = codes.len();
Ok(Self {
n,
codes,
occupancy,
capacity,
distinct_survivors,
})
}
pub fn validate(&self) -> Result<(), TopologyError> {
if self.codes.is_empty() {
return Err(TopologyError::EmptyCodebook);
}
if self.occupancy.len() != self.codes.len() {
return Err(TopologyError::BadConfig {
field: "occupancy",
expected: "one entry per code",
found: format!(
"{} entries for {} codes",
self.occupancy.len(),
self.codes.len()
),
});
}
for code in &self.codes {
if code.n() != self.n {
return Err(TopologyError::SizeMismatch {
expected: self.n,
found: code.n(),
});
}
}
if self.capacity < self.codes.len() {
return Err(TopologyError::BadConfig {
field: "capacity",
expected: "at least the number of codes",
found: format!("{} for {} codes", self.capacity, self.codes.len()),
});
}
Ok(())
}
pub fn team_size(&self) -> usize {
self.n
}
pub fn codes(&self) -> &[Topology] {
&self.codes
}
pub fn len(&self) -> usize {
self.codes.len()
}
pub fn is_empty(&self) -> bool {
self.codes.is_empty()
}
pub fn capacity(&self) -> usize {
self.capacity
}
pub fn used_codes(&self) -> usize {
self.occupancy.iter().filter(|&&o| o > 0).count()
}
pub fn distinct_survivors(&self) -> usize {
self.distinct_survivors
}
pub fn decode(&self, code: usize) -> Option<&Topology> {
self.codes.get(code)
}
pub fn encode(&self, topology: &Topology) -> Result<usize, TopologyError> {
if topology.n() != self.n {
return Err(TopologyError::SizeMismatch {
expected: self.n,
found: topology.n(),
});
}
nearest(&self.codes, topology)
}
}
fn nearest(codes: &[Topology], topology: &Topology) -> Result<usize, TopologyError> {
let mut best = None;
for (idx, code) in codes.iter().enumerate() {
let d = code.hamming(topology)?;
match best {
None => best = Some((idx, d)),
Some((_, bd)) if d < bd => best = Some((idx, d)),
_ => {}
}
}
best.map(|(idx, _)| idx).ok_or(TopologyError::EmptyCodebook)
}
fn k_medoids(
points: &[Topology],
k: usize,
max_iterations: usize,
) -> Result<Vec<Topology>, TopologyError> {
debug_assert!(points.len() > k, "caller handles the exact-index case");
let mut medoids: Vec<usize> = vec![0];
while medoids.len() < k {
let mut best: Option<(usize, usize)> = None;
for (i, p) in points.iter().enumerate() {
if medoids.contains(&i) {
continue;
}
let mut d_min = usize::MAX;
for &m in &medoids {
d_min = d_min.min(points[m].hamming(p)?);
}
if best.is_none_or(|(_, bd)| d_min > bd) {
best = Some((i, d_min));
}
}
match best {
Some((i, _)) => medoids.push(i),
None => break,
}
}
medoids.sort_unstable();
for _ in 0..max_iterations {
let mut clusters: Vec<Vec<usize>> = vec![Vec::new(); medoids.len()];
for (i, p) in points.iter().enumerate() {
let mut best: Option<(usize, usize)> = None;
for (slot, &m) in medoids.iter().enumerate() {
let d = points[m].hamming(p)?;
if best.is_none_or(|(_, bd)| d < bd) {
best = Some((slot, d));
}
}
clusters[best.expect("at least one medoid").0].push(i);
}
let mut next = medoids.clone();
for (slot, members) in clusters.iter().enumerate() {
if members.is_empty() {
continue;
}
let mut best: Option<(usize, usize)> = None;
for &cand in members {
let mut total = 0usize;
for &other in members {
total += points[cand].hamming(&points[other])?;
}
if best.is_none_or(|(_, bt)| total < bt) {
best = Some((cand, total));
}
}
next[slot] = best.expect("non-empty cluster").0;
}
next.sort_unstable();
next.dedup();
if next == medoids {
break;
}
medoids = next;
}
Ok(medoids.into_iter().map(|i| points[i].clone()).collect())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::record::ExecutionRecord;
use crate::topology::CoordinationShape;
fn set(records: Vec<ExecutionRecord>) -> RecordSet {
RecordSet::new(records).unwrap()
}
fn rec(task: &str, t: Topology, u: f32, tokens: u64) -> ExecutionRecord {
ExecutionRecord::new(task, vec![0.5, 0.5], t, u, tokens)
}
fn six_family_records(n: usize) -> Vec<ExecutionRecord> {
let mut out = Vec::new();
for (ti, task) in ["t1", "t2", "t3"].iter().enumerate() {
for (fi, topo) in Topology::collection_protocol(n)
.unwrap()
.into_iter()
.enumerate()
{
out.push(rec(task, topo, 1.0, 100 + (ti * 10 + fi) as u64));
}
}
out
}
#[test]
fn capacity_beyond_the_survivor_count_stays_idle() {
let records = set(six_family_records(4));
let mut used = Vec::new();
for capacity in [8, 16, 32, 64] {
let book = Codebook::fit(
&records,
&CodebookConfig {
capacity,
..Default::default()
},
)
.unwrap();
used.push(book.used_codes());
}
assert_eq!(used, vec![6, 6, 6, 6]);
}
#[test]
fn only_reward_surviving_topologies_enter_the_codebook() {
let n = 4;
let records = set(vec![
rec("t", Topology::complete(n).unwrap(), 1.0, 100),
rec("t", Topology::chain(n).unwrap(), 0.0, 100),
rec("t", Topology::star(n, 0).unwrap(), 1.0, 100),
]);
let book = Codebook::fit(&records, &CodebookConfig::default()).unwrap();
assert_eq!(book.len(), 2);
assert!(book
.codes()
.iter()
.all(|c| c != &Topology::chain(n).unwrap()));
}
#[test]
fn a_set_with_no_survivors_is_an_error_not_an_empty_book() {
let records = set(vec![rec("t", Topology::chain(4).unwrap(), 0.0, 100)]);
assert!(matches!(
Codebook::fit(&records, &CodebookConfig::default()),
Err(TopologyError::NoRecords { kind: "surviving" })
));
}
#[test]
fn fitting_is_deterministic_under_record_shuffling() {
let mut a = six_family_records(4);
let book_a = Codebook::fit(&set(a.clone()), &CodebookConfig::default()).unwrap();
a.reverse();
let book_b = Codebook::fit(&set(a), &CodebookConfig::default()).unwrap();
assert_eq!(book_a.codes(), book_b.codes());
}
#[test]
fn quantization_compresses_when_survivors_exceed_capacity() {
let n = 5;
let mut records = Vec::new();
for seed in 0..12u64 {
records.push(rec(
"t",
Topology::erdos_renyi(n, 0.5, seed + 1).unwrap(),
1.0,
100,
));
}
let book = Codebook::fit(
&set(records),
&CodebookConfig {
capacity: 4,
..Default::default()
},
)
.unwrap();
assert!(book.len() <= 4, "got {}", book.len());
assert_eq!(book.distinct_survivors(), 12);
assert_eq!(book.occupancy.iter().sum::<usize>(), 12);
}
#[test]
fn every_code_decodes_to_a_topology_that_was_executed() {
let n = 5;
let mut records = Vec::new();
let mut executed = Vec::new();
for seed in 0..12u64 {
let t = Topology::erdos_renyi(n, 0.5, seed + 1).unwrap();
executed.push(t.clone());
records.push(rec("t", t, 1.0, 100));
}
let book = Codebook::fit(
&set(records),
&CodebookConfig {
capacity: 4,
..Default::default()
},
)
.unwrap();
for code in book.codes() {
assert!(
executed.iter().any(|e| e == code),
"code {} was never executed",
code.key()
);
}
}
#[test]
fn encode_finds_the_nearest_code() {
let n = 4;
let book = Codebook::from_topologies(vec![
Topology::empty(n).unwrap(),
Topology::complete(n).unwrap(),
])
.unwrap();
let mut nearly_complete = Topology::complete(n).unwrap();
nearly_complete.set_edge(0, 1, false).unwrap();
assert_eq!(book.encode(&nearly_complete).unwrap(), 1);
assert_eq!(book.encode(&Topology::chain(n).unwrap()).unwrap(), 0);
}
#[test]
fn encode_rejects_a_mismatched_team_size() {
let book = Codebook::from_topologies(vec![Topology::complete(4).unwrap()]).unwrap();
assert!(matches!(
book.encode(&Topology::complete(5).unwrap()),
Err(TopologyError::SizeMismatch { .. })
));
}
#[test]
fn cold_start_book_indexes_the_coordination_shapes() {
let n = 4;
let shapes: Vec<Topology> = CoordinationShape::ALL
.iter()
.map(|s| s.topology(n).unwrap())
.collect();
let book = Codebook::from_topologies(shapes).unwrap();
assert_eq!(book.len(), 4);
}
#[test]
fn zero_capacity_is_rejected() {
let records = set(six_family_records(4));
assert!(matches!(
Codebook::fit(
&records,
&CodebookConfig {
capacity: 0,
..Default::default()
}
),
Err(TopologyError::BadConfig { .. })
));
}
}