use serde::{Deserialize, Serialize};
use crate::codebook::{Codebook, CodebookConfig};
use crate::error::TopologyError;
use crate::predictor::{CodePredictor, PredictorConfig};
use crate::proxy::{ExecutionProxy, ProxyConfig, ProxyScore};
use crate::record::{RecordSet, DEFAULT_COST_WEIGHT};
use crate::topology::{shape_of, CoordinationShape, Topology};
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SelectorConfig {
pub codebook: CodebookConfig,
pub predictor: PredictorConfig,
pub proxy: ProxyConfig,
pub top_m: usize,
pub cost_weight: f32,
}
impl Default for SelectorConfig {
fn default() -> Self {
Self {
codebook: CodebookConfig::default(),
predictor: PredictorConfig::default(),
proxy: ProxyConfig::default(),
top_m: 5,
cost_weight: DEFAULT_COST_WEIGHT,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Candidate {
pub code: usize,
pub topology: Topology,
pub shape: Option<CoordinationShape>,
pub prior: f32,
pub score: ProxyScore,
pub objective: f32,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Selection {
pub topology: Topology,
pub code: usize,
pub shape: Option<CoordinationShape>,
pub score: ProxyScore,
pub objective: f32,
pub considered: Vec<Candidate>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(try_from = "TopologySelectorWire")]
pub struct TopologySelector {
codebook: Codebook,
predictor: CodePredictor,
proxy: ExecutionProxy,
top_m: usize,
cost_weight: f32,
#[serde(default, skip_serializing_if = "Option::is_none")]
embedder: Option<String>,
}
#[derive(Deserialize)]
struct TopologySelectorWire {
codebook: Codebook,
predictor: CodePredictor,
proxy: ExecutionProxy,
top_m: usize,
cost_weight: f32,
#[serde(default)]
embedder: Option<String>,
}
impl TryFrom<TopologySelectorWire> for TopologySelector {
type Error = TopologyError;
fn try_from(wire: TopologySelectorWire) -> Result<Self, Self::Error> {
if wire.codebook.len() != wire.predictor.codes() {
return Err(TopologyError::BadConfig {
field: "predictor",
expected: "one output per codebook entry",
found: format!(
"{} codes vs {} outputs",
wire.codebook.len(),
wire.predictor.codes()
),
});
}
if wire.codebook.team_size() != wire.proxy.team_size() {
return Err(TopologyError::SizeMismatch {
expected: wire.codebook.team_size(),
found: wire.proxy.team_size(),
});
}
if wire.predictor.query_dim() != wire.proxy.query_dim() {
return Err(TopologyError::QueryDimMismatch {
expected: wire.predictor.query_dim(),
found: wire.proxy.query_dim(),
});
}
if wire.top_m == 0 {
return Err(TopologyError::BadConfig {
field: "top_m",
expected: "at least 1",
found: "0".into(),
});
}
if !wire.cost_weight.is_finite() {
return Err(TopologyError::BadConfig {
field: "cost_weight",
expected: "finite",
found: format!("{}", wire.cost_weight),
});
}
Ok(TopologySelector {
codebook: wire.codebook,
predictor: wire.predictor,
proxy: wire.proxy,
top_m: wire.top_m,
cost_weight: wire.cost_weight,
embedder: wire.embedder,
})
}
}
impl TopologySelector {
pub fn fit(records: &RecordSet, config: &SelectorConfig) -> Result<Self, TopologyError> {
if config.top_m == 0 {
return Err(TopologyError::BadConfig {
field: "top_m",
expected: "at least 1",
found: "0".into(),
});
}
if !config.cost_weight.is_finite() {
return Err(TopologyError::BadConfig {
field: "cost_weight",
expected: "finite",
found: format!("{}", config.cost_weight),
});
}
let codebook = Codebook::fit(records, &config.codebook)?;
let predictor = CodePredictor::fit(records, &codebook, &config.predictor)?;
let proxy = ExecutionProxy::fit(records, &config.proxy)?;
Ok(Self {
codebook,
predictor,
proxy,
top_m: config.top_m,
cost_weight: config.cost_weight,
embedder: records.embedder().map(str::to_owned),
})
}
pub fn from_parts(
codebook: Codebook,
predictor: CodePredictor,
proxy: ExecutionProxy,
top_m: usize,
cost_weight: f32,
) -> Result<Self, TopologyError> {
if codebook.len() != predictor.codes() {
return Err(TopologyError::BadConfig {
field: "predictor",
expected: "one output per codebook entry",
found: format!("{} codes vs {} outputs", codebook.len(), predictor.codes()),
});
}
if codebook.team_size() != proxy.team_size() {
return Err(TopologyError::SizeMismatch {
expected: codebook.team_size(),
found: proxy.team_size(),
});
}
if predictor.query_dim() != proxy.query_dim() {
return Err(TopologyError::QueryDimMismatch {
expected: predictor.query_dim(),
found: proxy.query_dim(),
});
}
if top_m == 0 {
return Err(TopologyError::BadConfig {
field: "top_m",
expected: "at least 1",
found: "0".into(),
});
}
if !cost_weight.is_finite() {
return Err(TopologyError::BadConfig {
field: "cost_weight",
expected: "finite",
found: format!("{cost_weight}"),
});
}
Ok(Self {
codebook,
predictor,
proxy,
top_m,
cost_weight,
embedder: None,
})
}
pub fn with_embedder(mut self, embedder: impl Into<String>) -> Self {
self.embedder = Some(embedder.into());
self
}
pub fn embedder(&self) -> Option<&str> {
self.embedder.as_deref()
}
pub fn codebook(&self) -> &Codebook {
&self.codebook
}
pub fn predictor(&self) -> &CodePredictor {
&self.predictor
}
pub fn proxy(&self) -> &ExecutionProxy {
&self.proxy
}
pub fn team_size(&self) -> usize {
self.codebook.team_size()
}
pub fn select_with(&self, query: &[f32], embedder: &str) -> Result<Selection, TopologyError> {
if let Some(fitted) = &self.embedder {
if fitted != embedder {
return Err(TopologyError::EmbedderMismatch {
fitted: fitted.clone(),
query: embedder.to_string(),
});
}
}
self.select(query)
}
pub fn select(&self, query: &[f32]) -> Result<Selection, TopologyError> {
let prior = self.predictor.predict(query)?;
let codes = self.predictor.top_codes(query, self.top_m)?;
let conditioned = self.proxy.condition(query)?;
let mut considered: Vec<Candidate> = Vec::with_capacity(codes.len());
for code in codes {
let topology = self
.codebook
.decode(code)
.ok_or(TopologyError::EmptyCodebook)?
.clone();
if considered.iter().any(|c| c.topology == topology) {
continue;
}
let score = conditioned.score(&topology, query)?;
considered.push(Candidate {
code,
shape: shape_of(&topology),
topology,
prior: prior.get(code).copied().unwrap_or(0.0),
score,
objective: score.objective(self.cost_weight),
});
}
let winner = considered
.iter()
.enumerate()
.max_by(|(ai, a), (bi, b)| {
a.objective
.total_cmp(&b.objective)
.then(b.code.cmp(&a.code))
.then(bi.cmp(ai))
})
.map(|(_, c)| c.clone())
.ok_or(TopologyError::EmptyCodebook)?;
Ok(Selection {
topology: winner.topology,
code: winner.code,
shape: winner.shape,
score: winner.score,
objective: winner.objective,
considered,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::record::ExecutionRecord;
fn records() -> RecordSet {
let n = 4;
let complete = Topology::complete(n).unwrap();
let chain = Topology::chain(n).unwrap();
let star = Topology::star(n, 0).unwrap();
let mut out = Vec::new();
for i in 0..6 {
let drift = i as f32 * 0.01;
let math_q = vec![1.0, 0.0, drift];
let math = format!("math{i}");
out.push(ExecutionRecord::new(
&math,
math_q.clone(),
complete.clone(),
1.0,
600,
));
out.push(ExecutionRecord::new(
&math,
math_q.clone(),
star.clone(),
1.0,
1200,
));
out.push(ExecutionRecord::new(
&math,
math_q,
chain.clone(),
1.0,
2400,
));
let code_q = vec![0.0, 1.0, drift];
let code = format!("code{i}");
out.push(ExecutionRecord::new(
&code,
code_q.clone(),
chain.clone(),
1.0,
600,
));
out.push(ExecutionRecord::new(
&code,
code_q.clone(),
star.clone(),
1.0,
1200,
));
out.push(ExecutionRecord::new(
&code,
code_q,
complete.clone(),
1.0,
2400,
));
}
RecordSet::new(out).unwrap()
}
fn selector() -> TopologySelector {
TopologySelector::fit(&records(), &SelectorConfig::default()).unwrap()
}
#[test]
fn selection_adapts_to_the_query() {
let s = selector();
let math = s.select(&[1.0, 0.0, 0.0]).unwrap();
let code = s.select(&[0.0, 1.0, 0.0]).unwrap();
assert_eq!(math.topology, Topology::complete(4).unwrap());
assert_eq!(code.topology, Topology::chain(4).unwrap());
assert_ne!(math.code, code.code);
}
#[test]
fn the_winner_maximizes_the_objective_among_the_candidates() {
let s = selector();
let selection = s.select(&[1.0, 0.0, 0.0]).unwrap();
let best = selection
.considered
.iter()
.map(|c| c.objective)
.fold(f32::NEG_INFINITY, f32::max);
assert!((selection.objective - best).abs() < 1e-6);
}
#[test]
fn selection_is_deterministic() {
let s = selector();
let a = s.select(&[0.3, 0.7, 0.1]).unwrap();
let b = s.select(&[0.3, 0.7, 0.1]).unwrap();
assert_eq!(a, b);
}
#[test]
fn candidates_are_deduplicated_and_bounded_by_top_m() {
let s = TopologySelector::fit(
&records(),
&SelectorConfig {
top_m: 2,
..Default::default()
},
)
.unwrap();
let selection = s.select(&[1.0, 0.0, 0.0]).unwrap();
assert!(selection.considered.len() <= 2);
for pair in 0..selection.considered.len() {
for other in (pair + 1)..selection.considered.len() {
assert_ne!(
selection.considered[pair].topology,
selection.considered[other].topology
);
}
}
}
#[test]
fn the_selection_reports_the_shape_a_caller_can_execute() {
let s = selector();
let selection = s.select(&[0.0, 1.0, 0.0]).unwrap();
assert_eq!(selection.shape, Some(CoordinationShape::Pipeline));
}
#[test]
fn reranking_beats_taking_the_prior_top_1_on_cost() {
let full = TopologySelector::fit(&records(), &SelectorConfig::default()).unwrap();
let top1 = TopologySelector::fit(
&records(),
&SelectorConfig {
top_m: 1,
..Default::default()
},
)
.unwrap();
let q = [1.0, 0.0, 0.0];
assert!(full.select(&q).unwrap().objective >= top1.select(&q).unwrap().objective - 1e-6);
}
#[test]
fn a_cold_start_selector_still_selects() {
let n = 4;
let codebook = Codebook::from_topologies(
CoordinationShape::ALL
.iter()
.map(|s| s.topology(n).unwrap())
.collect(),
)
.unwrap();
let predictor = CodePredictor::uniform(codebook.len(), 3).unwrap();
let proxy = ExecutionProxy::fit(&records(), &ProxyConfig::default()).unwrap();
let s = TopologySelector::from_parts(codebook, predictor, proxy, 5, DEFAULT_COST_WEIGHT)
.unwrap();
let selection = s.select(&[1.0, 0.0, 0.0]).unwrap();
assert!(selection.objective.is_finite());
assert!(!selection.considered.is_empty());
}
#[test]
fn mismatched_parts_are_rejected() {
let n = 4;
let codebook = Codebook::from_topologies(vec![Topology::complete(n).unwrap()]).unwrap();
let predictor = CodePredictor::uniform(3, 3).unwrap();
let proxy = ExecutionProxy::fit(&records(), &ProxyConfig::default()).unwrap();
assert!(matches!(
TopologySelector::from_parts(codebook, predictor, proxy, 5, 0.1),
Err(TopologyError::BadConfig {
field: "predictor",
..
})
));
}
#[test]
fn zero_top_m_is_rejected() {
assert!(matches!(
TopologySelector::fit(
&records(),
&SelectorConfig {
top_m: 0,
..Default::default()
}
),
Err(TopologyError::BadConfig { field: "top_m", .. })
));
}
#[test]
fn a_wrong_dimension_query_is_rejected_at_selection_time() {
let s = selector();
assert!(matches!(
s.select(&[1.0, 0.0]),
Err(TopologyError::QueryDimMismatch { .. })
));
}
#[test]
fn a_selector_round_trips_through_json() {
let s = selector();
let json = serde_json::to_string(&s).unwrap();
let back: TopologySelector = serde_json::from_str(&json).unwrap();
assert_eq!(s, back);
assert_eq!(
s.select(&[1.0, 0.0, 0.0]).unwrap(),
back.select(&[1.0, 0.0, 0.0]).unwrap()
);
}
}