use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use crate::codebook::Codebook;
use crate::error::TopologyError;
use crate::record::{RecordSet, DEFAULT_COST_WEIGHT};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PredictorConfig {
pub gamma: f32,
pub cost_weight: f32,
pub temperature: f32,
}
impl Default for PredictorConfig {
fn default() -> Self {
Self {
gamma: 2.0,
cost_weight: DEFAULT_COST_WEIGHT,
temperature: 8.0,
}
}
}
impl PredictorConfig {
fn validate(&self) -> Result<(), TopologyError> {
for (field, value) in [
("gamma", self.gamma),
("cost_weight", self.cost_weight),
("temperature", self.temperature),
] {
if !value.is_finite() {
return Err(TopologyError::BadConfig {
field,
expected: "finite",
found: format!("{value}"),
});
}
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
struct Condition {
direction: Vec<f32>,
has_direction: bool,
soft_target: Vec<f32>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(try_from = "CodePredictorWire")]
pub struct CodePredictor {
conditions: Vec<Condition>,
codes: usize,
query_dim: usize,
temperature: f32,
}
#[derive(Deserialize)]
struct CodePredictorWire {
conditions: Vec<Condition>,
codes: usize,
query_dim: usize,
temperature: f32,
}
impl TryFrom<CodePredictorWire> for CodePredictor {
type Error = TopologyError;
fn try_from(wire: CodePredictorWire) -> Result<Self, Self::Error> {
let predictor = CodePredictor {
conditions: wire.conditions,
codes: wire.codes,
query_dim: wire.query_dim,
temperature: wire.temperature,
};
predictor.validate()?;
Ok(predictor)
}
}
impl CodePredictor {
pub fn fit(
records: &RecordSet,
codebook: &Codebook,
config: &PredictorConfig,
) -> Result<Self, TopologyError> {
config.validate()?;
if codebook.is_empty() {
return Err(TopologyError::EmptyCodebook);
}
if codebook.team_size() != records.team_size() {
return Err(TopologyError::SizeMismatch {
expected: codebook.team_size(),
found: records.team_size(),
});
}
let k = codebook.len();
let mut order: Vec<&str> = Vec::new();
let mut groups: HashMap<&str, Vec<usize>> = HashMap::new();
for (index, record) in records.records().iter().enumerate() {
let group = groups.entry(record.task_id.as_str()).or_default();
if group.is_empty() {
order.push(record.task_id.as_str());
}
group.push(index);
}
let mut conditions = Vec::with_capacity(order.len());
for task_id in order {
let mut weights = vec![0f64; k];
let mut centroid = vec![0f32; records.query_dim()];
let mut count = 0usize;
for &index in &groups[task_id] {
let record = &records.records()[index];
let code = codebook.encode(&record.topology)?;
let reward = records.reward(index, config.cost_weight);
weights[code] += (config.gamma * reward).exp() as f64;
for (slot, value) in centroid.iter_mut().zip(record.query.iter()) {
*slot += value;
}
count += 1;
}
if count == 0 {
continue;
}
for slot in centroid.iter_mut() {
*slot /= count as f32;
}
let total: f64 = weights.iter().sum();
let soft_target: Vec<f32> = if total > 0.0 {
weights.iter().map(|w| (w / total) as f32).collect()
} else {
vec![1.0 / k as f32; k]
};
let (direction, has_direction) = normalize(¢roid);
conditions.push(Condition {
direction,
has_direction,
soft_target,
});
}
if conditions.is_empty() {
return Err(TopologyError::NoRecords { kind: "condition" });
}
Ok(Self {
conditions,
codes: k,
query_dim: records.query_dim(),
temperature: config.temperature,
})
}
pub fn uniform(codes: usize, query_dim: usize) -> Result<Self, TopologyError> {
if codes == 0 {
return Err(TopologyError::EmptyCodebook);
}
Ok(Self {
conditions: vec![Condition {
direction: vec![0.0; query_dim],
has_direction: false,
soft_target: vec![1.0 / codes as f32; codes],
}],
codes,
query_dim,
temperature: PredictorConfig::default().temperature,
})
}
pub fn validate(&self) -> Result<(), TopologyError> {
if self.codes == 0 {
return Err(TopologyError::EmptyCodebook);
}
if self.conditions.is_empty() {
return Err(TopologyError::NoRecords { kind: "condition" });
}
if !self.temperature.is_finite() {
return Err(TopologyError::BadConfig {
field: "temperature",
expected: "finite",
found: format!("{}", self.temperature),
});
}
for condition in &self.conditions {
if condition.soft_target.len() != self.codes {
return Err(TopologyError::BadConfig {
field: "soft_target",
expected: "one entry per code",
found: format!(
"{} entries for {} codes",
condition.soft_target.len(),
self.codes
),
});
}
if condition.direction.len() != self.query_dim {
return Err(TopologyError::QueryDimMismatch {
expected: self.query_dim,
found: condition.direction.len(),
});
}
}
Ok(())
}
pub fn codes(&self) -> usize {
self.codes
}
pub fn conditions(&self) -> usize {
self.conditions.len()
}
pub fn query_dim(&self) -> usize {
self.query_dim
}
pub fn predict(&self, query: &[f32]) -> Result<Vec<f32>, TopologyError> {
if query.len() != self.query_dim {
return Err(TopologyError::QueryDimMismatch {
expected: self.query_dim,
found: query.len(),
});
}
let (direction, has_direction) = normalize(query);
let logits: Vec<f32> = self
.conditions
.iter()
.map(|c| {
if has_direction && c.has_direction {
self.temperature * dot(&direction, &c.direction)
} else {
0.0
}
})
.collect();
let weights = softmax(&logits);
let mut out = vec![0f32; self.codes];
for (w, condition) in weights.iter().zip(self.conditions.iter()) {
for (slot, value) in out.iter_mut().zip(condition.soft_target.iter()) {
*slot += w * value;
}
}
let total: f32 = out.iter().sum();
if total > 0.0 {
for slot in out.iter_mut() {
*slot /= total;
}
} else {
out = vec![1.0 / self.codes as f32; self.codes];
}
Ok(out)
}
pub fn top_codes(&self, query: &[f32], m: usize) -> Result<Vec<usize>, TopologyError> {
let probabilities = self.predict(query)?;
let mut ranked: Vec<usize> = (0..probabilities.len()).collect();
ranked.sort_by(|&a, &b| {
probabilities[b]
.total_cmp(&probabilities[a])
.then(a.cmp(&b))
});
ranked.truncate(m.min(probabilities.len()));
Ok(ranked)
}
}
fn normalize(v: &[f32]) -> (Vec<f32>, bool) {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > f32::EPSILON {
(v.iter().map(|x| x / norm).collect(), true)
} else {
(vec![0.0; v.len()], false)
}
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b.iter()).map(|(x, y)| x * y).sum()
}
fn softmax(logits: &[f32]) -> Vec<f32> {
let max = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = logits.iter().map(|l| (l - max).exp()).collect();
let total: f32 = exps.iter().sum();
if total > 0.0 {
exps.into_iter().map(|e| e / total).collect()
} else {
vec![1.0 / logits.len() as f32; logits.len()]
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::codebook::CodebookConfig;
use crate::record::ExecutionRecord;
use crate::topology::{CoordinationShape, Topology};
fn split_records() -> RecordSet {
let n = 4;
let debate = CoordinationShape::Debate.topology(n).unwrap();
let pipeline = CoordinationShape::Pipeline.topology(n).unwrap();
let mut out = Vec::new();
for i in 0..4 {
let math_q = vec![1.0, 0.0, i as f32 * 0.01];
out.push(ExecutionRecord::new(
format!("math{i}"),
math_q.clone(),
debate.clone(),
1.0,
100,
));
out.push(ExecutionRecord::new(
format!("math{i}"),
math_q,
pipeline.clone(),
1.0,
900,
));
let code_q = vec![0.0, 1.0, i as f32 * 0.01];
out.push(ExecutionRecord::new(
format!("code{i}"),
code_q.clone(),
pipeline.clone(),
1.0,
100,
));
out.push(ExecutionRecord::new(
format!("code{i}"),
code_q,
debate.clone(),
1.0,
900,
));
}
RecordSet::new(out).unwrap()
}
fn fitted() -> (RecordSet, Codebook, CodePredictor) {
let records = split_records();
let book = Codebook::fit(&records, &CodebookConfig::default()).unwrap();
let predictor = CodePredictor::fit(&records, &book, &PredictorConfig::default()).unwrap();
(records, book, predictor)
}
#[test]
fn prediction_is_a_distribution() {
let (_, _, predictor) = fitted();
let p = predictor.predict(&[1.0, 0.0, 0.0]).unwrap();
assert_eq!(p.len(), predictor.codes());
assert!((p.iter().sum::<f32>() - 1.0).abs() < 1e-5);
assert!(p.iter().all(|v| (0.0..=1.0).contains(v)));
}
#[test]
fn the_prior_prefers_the_cheap_topology_for_each_query_family() {
let (_, book, predictor) = fitted();
let n = 4;
let debate = book
.encode(&CoordinationShape::Debate.topology(n).unwrap())
.unwrap();
let pipeline = book
.encode(&CoordinationShape::Pipeline.topology(n).unwrap())
.unwrap();
let math = predictor.predict(&[1.0, 0.0, 0.0]).unwrap();
assert!(
math[debate] > math[pipeline],
"math query should prefer debate: {math:?}"
);
let code = predictor.predict(&[0.0, 1.0, 0.0]).unwrap();
assert!(
code[pipeline] > code[debate],
"code query should prefer pipeline: {code:?}"
);
}
#[test]
fn utility_ties_are_broken_by_measured_cost() {
let (_, book, predictor) = fitted();
let n = 4;
let debate = book
.encode(&CoordinationShape::Debate.topology(n).unwrap())
.unwrap();
let p = predictor.predict(&[1.0, 0.0, 0.0]).unwrap();
assert!(p[debate] > 1.0 / book.len() as f32);
}
#[test]
fn top_codes_is_ranked_and_bounded() {
let (_, book, predictor) = fitted();
let top = predictor.top_codes(&[1.0, 0.0, 0.0], 1).unwrap();
assert_eq!(top.len(), 1);
let all = predictor.top_codes(&[1.0, 0.0, 0.0], 99).unwrap();
assert_eq!(all.len(), book.len());
let p = predictor.predict(&[1.0, 0.0, 0.0]).unwrap();
for pair in all.windows(2) {
assert!(p[pair[0]] >= p[pair[1]]);
}
}
#[test]
fn a_zero_query_gets_a_blend_not_a_nan() {
let (_, _, predictor) = fitted();
let p = predictor.predict(&[0.0, 0.0, 0.0]).unwrap();
assert!(p.iter().all(|v| v.is_finite()));
assert!((p.iter().sum::<f32>() - 1.0).abs() < 1e-5);
}
#[test]
fn wrong_query_dimension_is_rejected() {
let (_, _, predictor) = fitted();
assert!(matches!(
predictor.predict(&[1.0, 0.0]),
Err(TopologyError::QueryDimMismatch { .. })
));
}
#[test]
fn uniform_predictor_is_query_independent() {
let predictor = CodePredictor::uniform(4, 3).unwrap();
let a = predictor.predict(&[1.0, 0.0, 0.0]).unwrap();
let b = predictor.predict(&[0.0, 0.0, 1.0]).unwrap();
assert_eq!(a, b);
assert!((a[0] - 0.25).abs() < 1e-6);
}
#[test]
fn a_codebook_over_a_different_team_size_is_rejected() {
let records = split_records();
let book = Codebook::from_topologies(vec![Topology::complete(5).unwrap()]).unwrap();
assert!(matches!(
CodePredictor::fit(&records, &book, &PredictorConfig::default()),
Err(TopologyError::SizeMismatch { .. })
));
}
#[test]
fn fitting_stays_linear_as_the_task_count_grows() {
let n = 4;
let mut out = Vec::new();
for task in 0..2000 {
let q = task as f32 / 2000.0;
for topology in Topology::collection_protocol(n).unwrap() {
let edges = topology.edge_count() as u64;
out.push(ExecutionRecord::new(
format!("t{task}"),
vec![q, 1.0 - q],
topology,
1.0,
2400u64.saturating_sub(120 * edges),
));
}
}
let records = RecordSet::new(out).unwrap();
assert_eq!(records.len(), 12_000);
assert_eq!(records.task_ids().len(), 2000);
let book = Codebook::fit(&records, &CodebookConfig::default()).unwrap();
let predictor = CodePredictor::fit(&records, &book, &PredictorConfig::default()).unwrap();
assert_eq!(predictor.conditions(), 2000);
let p = predictor.predict(&[0.5, 0.5]).unwrap();
assert!((p.iter().sum::<f32>() - 1.0).abs() < 1e-4);
}
#[test]
fn grouping_preserves_first_seen_task_order() {
let n = 4;
let mut out = Vec::new();
for task in ["zebra", "apple", "mango"] {
for topology in [Topology::complete(n).unwrap(), Topology::chain(n).unwrap()] {
out.push(ExecutionRecord::new(
task,
vec![0.5, 0.5],
topology,
1.0,
900,
));
}
}
let records = RecordSet::new(out).unwrap();
assert_eq!(records.task_ids(), vec!["zebra", "apple", "mango"]);
let book = Codebook::fit(&records, &CodebookConfig::default()).unwrap();
let a = CodePredictor::fit(&records, &book, &PredictorConfig::default()).unwrap();
let b = CodePredictor::fit(&records, &book, &PredictorConfig::default()).unwrap();
assert_eq!(a, b, "grouping must not depend on hash iteration order");
}
#[test]
fn fitting_is_deterministic() {
let (records, book, _) = fitted();
let a = CodePredictor::fit(&records, &book, &PredictorConfig::default()).unwrap();
let b = CodePredictor::fit(&records, &book, &PredictorConfig::default()).unwrap();
assert_eq!(a, b);
}
}