use std::path::Path;
pub fn topic_cap_drops(assign: &[usize], sims: &[f32], k: usize, cap: usize) -> Vec<usize> {
debug_assert!(cap >= 1, "cap must be >= 1");
let mut members: Vec<Vec<usize>> = vec![Vec::new(); k];
for (i, &t) in assign.iter().enumerate() {
if t < k {
members[t].push(i);
}
}
let mut drops: Vec<usize> = Vec::new();
for m in members.iter_mut() {
if m.len() <= cap {
continue;
}
m.sort_by(|&a, &b| {
sims[b].partial_cmp(&sims[a]).unwrap_or(std::cmp::Ordering::Equal).then(a.cmp(&b))
});
drops.extend_from_slice(&m[cap..]);
}
drops.sort_unstable();
drops
}
pub struct RebalanceOutcome {
pub drops: Vec<usize>,
pub clusters: crate::cluster::ClusterResult,
}
pub async fn rebalance_inline(repo_root: &Path, train_answers: &[String]) -> std::io::Result<Option<RebalanceOutcome>> {
let cfg = crate::config::load_config(&repo_root.join(crate::config::CONFIG_FILE));
if !cfg.cluster.rebalance || train_answers.is_empty() || cfg.understand.embed.base_url.is_empty() {
return Ok(None);
}
let store_dir = repo_root.join(&cfg.understand.embed.store);
let vecs = match crate::embed::EndpointEmbedder::new(&cfg.understand.embed, cfg.network.proxy.as_deref()) {
Ok(e) => match crate::vectors::get_or_embed(&e, &store_dir, &cfg.understand.embed.model, train_answers, cfg.understand.embed.batch_size).await {
Ok(v) => v,
Err(err) => { eprintln!("kibble: rebalance skipped (embed failed): {err}"); return Ok(None); }
},
Err(err) => { eprintln!("kibble: rebalance skipped (embed init failed): {err}"); return Ok(None); }
};
if vecs.is_empty() {
return Ok(None);
}
let c = crate::cluster::cluster_from(train_answers, &vecs, cfg.cluster.k, &cfg.understand.embed.model, cfg.cluster.min_cluster_size);
let assign = crate::cluster::assign_topics(&c.centroids, &vecs);
let sims: Vec<f32> = assign.iter().enumerate()
.map(|(i, &t)| crate::vectors::cosine(&vecs[i], &c.centroids[t]))
.collect();
let cap = ((cfg.cluster.max_topic_share * train_answers.len() as f64).floor() as usize).max(1);
let drops = topic_cap_drops(&assign, &sims, c.centroids.len(), cap);
let drop_set: std::collections::HashSet<usize> = drops.iter().copied().collect();
let kk = c.centroids.len();
let mut sizes = vec![0usize; kk];
let mut first: Vec<Option<usize>> = vec![None; kk];
for (i, &t) in assign.iter().enumerate() {
if drop_set.contains(&i) { continue; }
sizes[t] += 1;
if first[t].is_none() { first[t] = Some(i); }
}
let samples: Vec<String> = (0..kk)
.map(|t| first[t].map(|i| train_answers[i].chars().take(80).collect::<String>()).unwrap_or_default())
.collect();
let mut clusters = crate::cluster::ClusterResult {
k: c.k, model: c.model, centroids: c.centroids, labels: c.labels, sizes, samples,
};
crate::cluster::apply_llm_labels(&cfg, &mut clusters).await;
Ok(Some(RebalanceOutcome { drops, clusters }))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn drops_lowest_sim_over_cap() {
let assign = vec![0, 0, 0, 1];
let sims = vec![0.9, 0.5, 0.7, 0.99];
assert_eq!(topic_cap_drops(&assign, &sims, 2, 2), vec![1]);
}
#[test]
fn multi_topic_only_over_cap_dropped() {
let assign = vec![0, 0, 0, 1, 1];
let sims = vec![0.8, 0.6, 0.9, 0.7, 0.65];
assert_eq!(topic_cap_drops(&assign, &sims, 2, 2), vec![1]);
}
#[test]
fn tie_break_keeps_lower_index() {
let assign = vec![0, 0, 0];
let sims = vec![0.5, 0.5, 0.5];
assert_eq!(topic_cap_drops(&assign, &sims, 1, 2), vec![2]);
}
#[test]
fn no_drops_when_under_cap() {
let assign = vec![0, 1, 0, 1];
let sims = vec![0.5, 0.5, 0.5, 0.5];
assert!(topic_cap_drops(&assign, &sims, 2, 10).is_empty());
}
#[test]
fn inline_pipeline_caps_dominant_topic() {
let answers: Vec<String> = (0..7).map(|i| format!("doc {i}")).collect();
let vecs = vec![
vec![1.0, 0.0], vec![0.98, 0.02], vec![0.95, 0.05], vec![0.9, 0.1], vec![0.85, 0.15],
vec![0.0, 1.0], vec![0.05, 0.95],
];
let c = crate::cluster::cluster_from(&answers, &vecs, 2, "m", 0);
let assign = crate::cluster::assign_topics(&c.centroids, &vecs);
let sims: Vec<f32> = assign.iter().enumerate()
.map(|(i, &t)| crate::vectors::cosine(&vecs[i], &c.centroids[t])).collect();
let cap = 3usize;
let drops = topic_cap_drops(&assign, &sims, c.centroids.len(), cap);
assert_eq!(drops.len(), 2);
let dropset: std::collections::HashSet<usize> = drops.iter().copied().collect();
let mut counts = vec![0usize; c.centroids.len()];
for (i, &t) in assign.iter().enumerate() { if !dropset.contains(&i) { counts[t] += 1; } }
assert!(counts.iter().all(|&n| n <= cap), "counts {counts:?} exceed cap {cap}");
}
#[tokio::test]
async fn rebalance_inline_failsoft() {
let dir = std::env::temp_dir().join(format!("kibble_rbi_{}", std::process::id()));
let _ = std::fs::remove_dir_all(&dir);
std::fs::create_dir_all(&dir).unwrap();
let ans = vec!["a".to_string(), "b".to_string()];
std::fs::write(dir.join(crate::config::CONFIG_FILE), "[understand.embed]\nbase_url=\"http://127.0.0.1:1/v1\"\n").unwrap();
assert!(rebalance_inline(&dir, &ans).await.unwrap().is_none());
std::fs::write(dir.join(crate::config::CONFIG_FILE), "[cluster]\nrebalance=true\n[understand.embed]\nbase_url=\"http://127.0.0.1:1/v1\"\nmodel=\"nomic\"\n").unwrap();
assert!(rebalance_inline(&dir, &ans).await.unwrap().is_none());
std::fs::write(dir.join(crate::config::CONFIG_FILE), "[cluster]\nrebalance=true\n[understand.embed]\nbase_url=\"\"\n").unwrap();
assert!(rebalance_inline(&dir, &ans).await.unwrap().is_none());
std::fs::remove_dir_all(&dir).ok();
}
}