use std::{cmp::Ordering, collections::HashMap};
use crate::{
config,
superfile::vector::{
cell_posting::{
EncodedCellRow, dequantize_sq8_residual_into, manifest_centroid_components_from_row,
},
distance::{Metric, distance, nearest_k_centroids_transposed, relative_score_window},
kmeans::{kmeans, kmeans_pp},
},
supertable::manifest::{
ClusterCentroids, RABITQ_ADMIT_CELL_SHORTLIST_FRACTION, RABITQ_ADMIT_CELL_SHORTLIST_MIN,
RabitqAdmitContext,
},
};
pub(crate) fn cell_split_doc_cap() -> u64 {
config::global().vector.cell_split_doc_cap
}
pub(crate) fn split_overflow_needed(n_docs: u64) -> bool {
n_docs > cell_split_doc_cap()
}
pub(crate) fn cell_split_modality_d() -> f64 {
config::global().vector.cell_split_modality_d
}
const MODALITY_MODES_PER_CELL: usize = 4;
pub(crate) const MODALITY_MIN_CELL_DOCS: u64 = 128;
pub(crate) fn split_candidate(n_docs: u64) -> bool {
split_overflow_needed(n_docs)
|| (cell_split_modality_d() > 0.0 && n_docs >= MODALITY_MIN_CELL_DOCS)
}
pub(crate) fn apply_cell_count_updates(
base: &ClusterCentroids,
count_updates: &HashMap<u32, u32>,
) -> ClusterCentroids {
let mut updated = base.clone();
for (&cell, &count) in count_updates {
if let Some(slot) = updated.counts.get_mut(cell as usize) {
*slot = count;
}
}
updated
}
pub(crate) fn apply_cell_updates(
base: &ClusterCentroids,
count_updates: &HashMap<u32, u32>,
) -> ClusterCentroids {
apply_cell_count_updates(base, count_updates)
}
pub(crate) const REPLICA_CLOSURE_MAX_REPLICAS: usize = 3;
pub(crate) const REPLICA_CLOSURE_DISTANCE_RATIO: f32 = 1.2;
const SPLIT_KMEANS_SAMPLE_PER_CLUSTER: usize = 2048;
const SPLIT_KMEANS_SAMPLE_MIN: usize = 4096;
const SPLIT_KMEANS_ITERS: usize = 10;
const SPLIT_KMEANS_SEED_XOR: u64 = 0x5157_5f4b_4d45_414e;
const SPLIT_ROUTE_FIDELITY_TARGET: f64 = 0.97;
const SPLIT_SELF_TUNE_K_STEP: f64 = 1.5;
const SPLIT_SELF_TUNE_K_MAX_FACTOR: usize = 4;
#[derive(Debug, Clone, Copy, PartialEq)]
pub(crate) struct BoundaryAssignment {
pub primary: u32,
pub replicas: [Option<(u32, f32)>; REPLICA_CLOSURE_MAX_REPLICAS],
}
fn boundary_margin(
clusters: &ClusterCentroids,
metric: Metric,
primary: u32,
neighbor: u32,
primary_score: f32,
neighbor_score: f32,
) -> f32 {
let gap = (neighbor_score - primary_score).max(0.0);
let c1 = clusters.centroid(primary as usize);
let c2 = clusters.centroid(neighbor as usize);
match metric {
Metric::L2Sq => {
let separation = distance(metric, c1, c2).sqrt();
if separation > 0.0 {
gap / (2.0 * separation)
} else {
f32::INFINITY
}
}
Metric::Cosine | Metric::NegDot => {
let separation = distance(metric, c1, c2).abs();
if separation > 0.0 {
gap / separation
} else {
f32::INFINITY
}
}
}
}
pub(crate) fn assignment_shortlist_window(n_cells: usize) -> usize {
let scaled = (n_cells as f64 * RABITQ_ADMIT_CELL_SHORTLIST_FRACTION).ceil() as usize;
scaled
.max(RABITQ_ADMIT_CELL_SHORTLIST_MIN)
.min(n_cells.max(1))
}
pub(crate) fn boundary_assignment_encoded(
clusters: &ClusterCentroids,
metric: Metric,
row: &EncodedCellRow,
admit_ctx: &RabitqAdmitContext,
window: usize,
) -> BoundaryAssignment {
let dim = clusters.dim as usize;
let mut row_fp = vec![0f32; dim];
dequantize_sq8_residual_into(
&row.scale,
&row.offset,
&row.codes,
&row.residuals,
row.rerank_codec
.residual_divisor()
.expect("encoded row uses residual-family codec"),
&mut row_fp,
);
boundary_assignment_fp32(clusters, metric, &row_fp, admit_ctx, window)
}
pub(crate) fn boundary_assignment_fp32(
clusters: &ClusterCentroids,
metric: Metric,
row_fp: &[f32],
admit_ctx: &RabitqAdmitContext,
window: usize,
) -> BoundaryAssignment {
let n_cent = clusters.n_cent as usize;
let top_k = REPLICA_CLOSURE_MAX_REPLICAS + 1;
let ranked: Vec<(u32, f32)> = if window >= n_cent {
nearest_k_centroids_transposed(
metric,
row_fp,
clusters.transposed(),
n_cent,
clusters.dim as usize,
None,
top_k,
)
} else {
let admit = admit_ctx.encode(row_fp);
let mut exact: Vec<(u32, f32)> = clusters
.admit_shortlist(metric, &admit, window)
.into_iter()
.map(|(cell, _)| (cell, clusters.score_one(metric, cell as usize, row_fp)))
.collect();
exact.sort_unstable_by(|a, b| a.1.total_cmp(&b.1).then_with(|| a.0.cmp(&b.0)));
exact.truncate(top_k);
exact
};
boundary_from_ranked(clusters, metric, &ranked)
}
fn boundary_from_ranked(
clusters: &ClusterCentroids,
metric: Metric,
ranked: &[(u32, f32)],
) -> BoundaryAssignment {
let mut replicas = [None; REPLICA_CLOSURE_MAX_REPLICAS];
let Some(&(primary, primary_score)) = ranked.first() else {
return BoundaryAssignment {
primary: 0,
replicas,
};
};
let closure_threshold =
relative_score_window(primary_score, REPLICA_CLOSURE_DISTANCE_RATIO - 1.0);
for (slot, &(cell, score)) in ranked.iter().skip(1).enumerate() {
if score > closure_threshold {
break;
}
replicas[slot] = Some((
cell,
boundary_margin(clusters, metric, primary, cell, primary_score, score),
));
}
BoundaryAssignment { primary, replicas }
}
fn dequantize_row(row: &EncodedCellRow, dim: usize) -> Vec<f32> {
let mut out = vec![0f32; dim];
dequantize_sq8_residual_into(
&row.scale,
&row.offset,
&row.codes,
&row.residuals,
row.rerank_codec
.residual_divisor()
.expect("encoded row uses residual-family codec"),
&mut out,
);
out
}
fn ashman_d(points: &[f32], dim: usize, cents: &[f32]) -> f64 {
let m = points.len() / dim;
if m < 2 || dim == 0 || cents.len() < 2 * dim {
return 0.0;
}
let mut axis = vec![0f32; dim];
let mut norm2 = 0f64;
for j in 0..dim {
let d = cents[dim + j] - cents[j];
axis[j] = d;
norm2 += f64::from(d) * f64::from(d);
}
if norm2 <= 1e-12 {
return 0.0;
}
let project =
|v: &[f32]| -> f64 { (0..dim).map(|j| f64::from(v[j]) * f64::from(axis[j])).sum() };
let mid = 0.5 * (project(¢s[..dim]) + project(¢s[dim..2 * dim]));
let mut cnt = [0f64; 2];
let mut sum = [0f64; 2];
let mut sumsq = [0f64; 2];
for i in 0..m {
let p = project(&points[i * dim..(i + 1) * dim]);
let s = usize::from(p >= mid);
cnt[s] += 1.0;
sum[s] += p;
sumsq[s] += p * p;
}
if cnt[0] < 1.0 || cnt[1] < 1.0 {
return 0.0;
}
let mean = [sum[0] / cnt[0], sum[1] / cnt[1]];
let var = [
(sumsq[0] / cnt[0] - mean[0] * mean[0]).max(0.0),
(sumsq[1] / cnt[1] - mean[1] * mean[1]).max(0.0),
];
let denom = (var[0] + var[1]).sqrt();
if denom <= 0.0 {
return f64::INFINITY;
}
std::f64::consts::SQRT_2 * (mean[1] - mean[0]).abs() / denom
}
const MODALITY_MAX_DEPTH: usize = 6;
const MODALITY_RECURSE_SEED_LEFT: u64 = 0x1111;
const MODALITY_RECURSE_SEED_RIGHT: u64 = 0x2222;
fn decode_rows(rows: &[&EncodedCellRow], dim: usize) -> Vec<f32> {
let mut out = Vec::with_capacity(rows.len() * dim);
for &row in rows {
out.extend_from_slice(&dequantize_row(row, dim));
}
out
}
fn recursive_binary_k(
decoded: &[f32],
dim: usize,
idx: &[usize],
seed: u64,
threshold: f64,
depth: usize,
) -> usize {
let m = idx.len();
if (m as u64) < MODALITY_MIN_CELL_DOCS || depth == 0 {
return 1;
}
let sample_n = m.min((2 * SPLIT_KMEANS_SAMPLE_PER_CLUSTER).max(SPLIT_KMEANS_SAMPLE_MIN));
let mut sample = Vec::with_capacity(sample_n * dim);
for s in 0..sample_n {
let i = idx[s * m / sample_n];
sample.extend_from_slice(&decoded[i * dim..(i + 1) * dim]);
}
let cents = kmeans(&sample, dim, 2, SPLIT_KMEANS_ITERS, seed);
if cents.len() < 2 * dim || ashman_d(&sample, dim, ¢s) < threshold {
return 1;
}
let (c0, c1) = (¢s[..dim], ¢s[dim..2 * dim]);
let mut left = Vec::new();
let mut right = Vec::new();
for &i in idx {
let v = &decoded[i * dim..(i + 1) * dim];
if distance(Metric::L2Sq, v, c0) <= distance(Metric::L2Sq, v, c1) {
left.push(i);
} else {
right.push(i);
}
}
if left.is_empty() || right.is_empty() {
return 1;
}
let seed_l = seed ^ MODALITY_RECURSE_SEED_LEFT;
let seed_r = seed ^ MODALITY_RECURSE_SEED_RIGHT;
recursive_binary_k(decoded, dim, &left, seed_l, threshold, depth - 1)
+ recursive_binary_k(decoded, dim, &right, seed_r, threshold, depth - 1)
}
pub(crate) fn cell_split_plan(
rows: &[&EncodedCellRow],
dim: usize,
split_cell: u32,
modality_d: f64,
) -> Option<(usize, bool)> {
let n_docs = rows.len() as u64;
let cap = cell_split_doc_cap().max(1) as usize;
let k_by_cap = rows.len().div_ceil(cap).max(2);
let threshold = modality_d;
if threshold <= 0.0 {
return Some((k_by_cap, true));
}
if split_overflow_needed(n_docs) {
return Some((k_by_cap, true));
}
if n_docs < MODALITY_MIN_CELL_DOCS {
return None;
}
let seed = (split_cell as u64) ^ SPLIT_KMEANS_SEED_XOR;
let decoded = decode_rows(rows, dim);
let idx: Vec<usize> = (0..rows.len()).collect();
let k = recursive_binary_k(&decoded, dim, &idx, seed, threshold, MODALITY_MAX_DEPTH);
let r = MODALITY_MODES_PER_CELL;
if k <= r {
return None;
}
let g = k.div_ceil(r).max(2);
Some((g, true))
}
fn capacitated_split_at_k(
rows: &[&EncodedCellRow],
split_cell: u32,
dim: usize,
metric: Metric,
k: usize,
cap_target: usize,
) -> (Vec<f32>, Vec<u32>, f64) {
let n = rows.len();
let mut assign = vec![0u32; n];
let sample_n = n.min((k * SPLIT_KMEANS_SAMPLE_PER_CLUSTER).max(SPLIT_KMEANS_SAMPLE_MIN));
let mut sample = Vec::with_capacity(sample_n * dim);
for s in 0..sample_n {
let idx = s * n / sample_n;
sample.extend_from_slice(&dequantize_row(rows[idx], dim));
}
let seed = (split_cell as u64) ^ SPLIT_KMEANS_SEED_XOR;
let cents = if k > 2 {
kmeans_pp(&sample, dim, k, SPLIT_KMEANS_ITERS, seed)
} else {
kmeans(&sample, dim, k, SPLIT_KMEANS_ITERS, seed)
};
if cents.len() < k * dim {
return (cents, assign, 0.0);
}
let mut row_dists = vec![0f32; n * k];
let mut nearest = vec![0u32; n];
let mut order: Vec<(usize, f32)> = Vec::with_capacity(n);
for (i, row) in rows.iter().copied().enumerate() {
let rv = dequantize_row(row, dim);
let base = i * k;
let (mut best, mut second, mut best_c) = (f32::INFINITY, f32::INFINITY, 0usize);
for c in 0..k {
let d = distance(metric, &rv, ¢s[c * dim..(c + 1) * dim]);
row_dists[base + c] = d;
if d < best {
second = best;
best = d;
best_c = c;
} else if d < second {
second = d;
}
}
nearest[i] = best_c as u32;
order.push((i, second - best));
}
order.sort_unstable_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(Ordering::Equal));
let mut counts = vec![0usize; k];
for (i, _) in order {
let base = i * k;
let mut best_c = usize::MAX;
let mut best_d = f32::INFINITY;
for c in 0..k {
if counts[c] < cap_target && row_dists[base + c] < best_d {
best_d = row_dists[base + c];
best_c = c;
}
}
if best_c == usize::MAX {
best_c = (0..k)
.min_by(|&a, &b| {
row_dists[base + a]
.partial_cmp(&row_dists[base + b])
.unwrap_or(Ordering::Equal)
})
.unwrap_or(0);
}
assign[i] = best_c as u32;
counts[best_c] += 1;
}
let faithful = (0..n).filter(|&i| assign[i] == nearest[i]).count();
(cents, assign, faithful as f64 / n as f64)
}
pub(crate) fn plan_sq8_split_kway(
rows: &[&EncodedCellRow],
clusters: &ClusterCentroids,
split_cell: u32,
metric: Metric,
k: usize,
self_tune: bool,
) -> (Vec<f32>, Vec<u32>) {
let dim = clusters.dim as usize;
let k = k.max(2).min(rows.len().max(2));
if rows.len() < 2 {
let c = manifest_centroid_components_from_row(rows[0], dim);
let mut cents = Vec::with_capacity(k * dim);
for _ in 0..k {
cents.extend_from_slice(&c);
}
return (cents, vec![0u32; rows.len()]);
}
let n = rows.len();
let cap_target = n.div_ceil(k).max(1);
let k_max = if self_tune {
k.saturating_mul(SPLIT_SELF_TUNE_K_MAX_FACTOR).min(n).max(k)
} else {
k
};
let mut best: Option<(f64, Vec<f32>, Vec<u32>)> = None;
let mut k_try = k;
loop {
let (cents, cand, rf) =
capacitated_split_at_k(rows, split_cell, dim, metric, k_try, cap_target);
if best.as_ref().is_none_or(|b| rf > b.0) {
best = Some((rf, cents, cand));
}
if rf >= SPLIT_ROUTE_FIDELITY_TARGET || k_try >= k_max {
break;
}
k_try = ((k_try as f64 * SPLIT_SELF_TUNE_K_STEP).ceil() as usize)
.max(k_try + 1)
.min(k_max);
}
let (rf, cents, cand) = best.expect("self-tune loop sets best on the first iteration");
if tracing::enabled!(tracing::Level::DEBUG) {
let k_final = (cents.len() / dim).max(1);
let mut sizes = vec![0usize; k_final];
for &c in &cand {
sizes[c as usize] += 1;
}
tracing::debug!(
cell = split_cell,
rows = n,
cap_target,
k_start = k,
k_final,
route_fidelity = rf,
child_min = sizes.iter().copied().min().unwrap_or(0),
child_max = sizes.iter().copied().max().unwrap_or(0),
"cell split planned"
);
}
(cents, cand)
}
#[cfg(test)]
pub(crate) fn plan_sq8_split(
rows: &[&EncodedCellRow],
clusters: &ClusterCentroids,
split_cell: u32,
metric: Metric,
) -> (Vec<f32>, Vec<f32>, Vec<u8>) {
let dim = clusters.dim as usize;
let (cents, assign) = plan_sq8_split_kway(rows, clusters, split_cell, metric, 2, true);
let c0 = cents[..dim].to_vec();
let c1 = cents[dim..2 * dim].to_vec();
(c0, c1, assign.iter().map(|&a| a as u8).collect())
}
pub(crate) fn insert_split_centroids(
base: &ClusterCentroids,
cell_id: u32,
sub_centroids: &[f32],
k: usize,
) -> (ClusterCentroids, Vec<u32>) {
let dim = base.dim as usize;
let p = cell_id as usize;
let old_n = base.n_cent as usize;
let new_n = old_n + (k - 1);
let mut fp32 = vec![0f32; new_n * dim];
for c in 0..old_n {
fp32[c * dim..(c + 1) * dim].copy_from_slice(base.centroid(c));
}
fp32[p * dim..(p + 1) * dim].copy_from_slice(&sub_centroids[..dim]);
let mut ids = vec![cell_id];
for j in 1..k {
let new_id = old_n + (j - 1);
fp32[new_id * dim..(new_id + 1) * dim]
.copy_from_slice(&sub_centroids[j * dim..(j + 1) * dim]);
ids.push(new_id as u32);
}
let mut counts = base.counts.clone();
counts.resize(new_n, 0);
let updated = ClusterCentroids::from_fp32(new_n as u32, base.dim, &fp32, counts);
(updated, ids)
}
#[cfg(test)]
pub(crate) fn insert_split_centroid(
base: &ClusterCentroids,
cell_id: u32,
sub_centroids: &[f32],
) -> (ClusterCentroids, u32) {
let (updated, ids) = insert_split_centroids(base, cell_id, sub_centroids, 2);
(updated, ids[1])
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::superfile::vector::{
cell_posting::{encode_blob, load_encoded_rows_from_blob},
rerank_codec::{RerankCodec, SQ8_FIXED_OFFSET, SQ8_FIXED_SCALE},
};
fn synth_centroids(n_cent: u32, dim: u32) -> ClusterCentroids {
let nc = n_cent as usize;
let d = dim as usize;
let mut fp32 = vec![0f32; nc * d];
for c in 0..nc {
for j in 0..d {
fp32[c * d + j] = c as f32 * 0.5 + j as f32 * 0.01;
}
}
let counts = vec![100; nc];
ClusterCentroids::from_fp32(n_cent, dim, &fp32, counts)
}
fn synth_rows(dim: usize, n: usize, offset: f32) -> Vec<EncodedCellRow> {
let mut ids = Vec::new();
let mut vecs = Vec::new();
for i in 0..n as u32 {
ids.push(i);
for d in 0..dim {
vecs.push(offset + i as f32 * 0.01 + d as f32 * 0.001);
}
}
let blob =
encode_blob(Metric::L2Sq, dim, &ids, &vecs, RerankCodec::Sq8Residual).expect("encode");
let stable_ids: Vec<i128> = (0..n).map(|i| i as i128).collect();
load_encoded_rows_from_blob(&blob, &stable_ids, None).expect("load")
}
fn synth_gaussian_cell(
dim: usize,
n_blobs: usize,
per_blob: usize,
sigma: f32,
seed: u64,
) -> Vec<EncodedCellRow> {
use rand::{SeedableRng, rngs::StdRng};
use rand_distr::{Distribution, Normal};
let mut rng = StdRng::seed_from_u64(seed);
let unit = Normal::new(0.0f32, 1.0).expect("unit normal");
let noise = Normal::new(0.0f32, sigma).expect("noise normal");
let n = n_blobs * per_blob;
let mut ids = Vec::with_capacity(n);
let mut vecs = Vec::with_capacity(n * dim);
for _ in 0..n_blobs {
let center: Vec<f32> = (0..dim).map(|_| unit.sample(&mut rng)).collect();
for _ in 0..per_blob {
let mut v: Vec<f32> = (0..dim)
.map(|d| center[d] + noise.sample(&mut rng))
.collect();
crate::superfile::vector::distance::normalize(&mut v);
ids.push(ids.len() as u32);
vecs.extend_from_slice(&v);
}
}
let blob =
encode_blob(Metric::L2Sq, dim, &ids, &vecs, RerankCodec::Sq8Residual).expect("encode");
let stable_ids: Vec<i128> = (0..n).map(|i| i as i128).collect();
load_encoded_rows_from_blob(&blob, &stable_ids, None).expect("load")
}
fn synth_fixed_rows(dim: usize, n: usize, code: u8) -> Vec<EncodedCellRow> {
let scale: Arc<[f32]> = Arc::from(vec![SQ8_FIXED_SCALE; dim]);
let offset: Arc<[f32]> = Arc::from(vec![SQ8_FIXED_OFFSET; dim]);
(0..n)
.map(|id| EncodedCellRow {
stable_id: id as i128,
rerank_codec: RerankCodec::Sq8FixedResidual,
scale: Arc::clone(&scale),
offset: Arc::clone(&offset),
codes: vec![code; dim],
residuals: vec![0; dim],
norm_sq: None,
})
.collect()
}
const TEST_ROT_SEED: u64 = 7;
#[test]
fn boundary_assignment_closure_matches_distance_ratio() {
let dim = 4usize;
let mut fp32 = Vec::new();
for base in [0.0f32, 1.0, 2.0, 30.0] {
fp32.extend(std::iter::repeat_n(base, dim));
}
let clusters = ClusterCentroids::from_fp32(4, dim as u32, &fp32, vec![1; 4]);
let ctx = RabitqAdmitContext::new(dim, TEST_ROT_SEED);
let window = assignment_shortlist_window(4);
let deep = vec![0.9f32; dim];
let assignment = boundary_assignment_fp32(&clusters, Metric::L2Sq, &deep, &ctx, window);
assert_eq!(assignment.primary, 1);
assert_eq!(assignment.replicas, [None; REPLICA_CLOSURE_MAX_REPLICAS]);
let boundary = vec![1.5f32; dim];
let assignment = boundary_assignment_fp32(&clusters, Metric::L2Sq, &boundary, &ctx, window);
assert_eq!(assignment.primary, 1);
assert_eq!(assignment.replicas[0].map(|(cell, _)| cell), Some(2));
assert_eq!(assignment.replicas[1], None);
let margin = assignment.replicas[0].expect("replica").1;
assert!(
margin.is_finite() && margin >= 0.0,
"boundary margin must be a finite non-negative distance, got {margin}"
);
}
#[test]
fn assignment_shortlist_window_scales_with_grid() {
assert_eq!(assignment_shortlist_window(1), 1);
assert_eq!(assignment_shortlist_window(16), 16);
assert_eq!(assignment_shortlist_window(48), 48);
assert_eq!(
assignment_shortlist_window(64),
RABITQ_ADMIT_CELL_SHORTLIST_MIN
);
assert_eq!(
assignment_shortlist_window(240),
RABITQ_ADMIT_CELL_SHORTLIST_MIN
);
assert_eq!(assignment_shortlist_window(256), 52);
assert_eq!(assignment_shortlist_window(512), 103);
assert_eq!(assignment_shortlist_window(1024), 205);
}
#[test]
fn shortlisted_assignment_matches_exact_on_planted_cells() {
let dim = 64usize;
let n_cells = 300usize;
let mut fp32 = vec![0.0f32; n_cells * dim];
for (c, chunk) in fp32.chunks_mut(dim).enumerate() {
chunk[c % dim] = 4.0 + (c / dim) as f32;
chunk[(c * 7 + 3) % dim] = 2.0;
}
let clusters =
ClusterCentroids::from_fp32(n_cells as u32, dim as u32, &fp32, vec![1; n_cells]);
let ctx = RabitqAdmitContext::new(dim, TEST_ROT_SEED);
let window = assignment_shortlist_window(n_cells);
assert!(window < n_cells, "test must exercise the shortlist arm");
let mut state = 0x9e37_79b9_97f4_a7c5u64;
let mut jitter = || {
state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1);
((state >> 33) % 1000) as f32 / 1000.0 * 0.2 - 0.1
};
for c in 0..n_cells {
let mut row = fp32[c * dim..(c + 1) * dim].to_vec();
for v in row.iter_mut() {
*v += jitter();
}
let shortlisted = boundary_assignment_fp32(&clusters, Metric::L2Sq, &row, &ctx, window);
let exact = boundary_assignment_fp32(&clusters, Metric::L2Sq, &row, &ctx, n_cells);
assert_eq!(
shortlisted.primary, exact.primary,
"cell {c}: shortlisted primary diverged from exact"
);
assert_eq!(shortlisted.primary, c as u32, "cell {c}: wrong placement");
}
}
#[test]
fn insert_split_centroid_extends_n_cent() {
let base = synth_centroids(4, 8);
let sub = vec![
0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 1.1, 1.2, 1.3, 1.4, 1.5, 1.6, 1.7, 1.8,
];
let (updated, new_id) = insert_split_centroid(&base, 2, &sub);
assert_eq!(new_id, 4);
assert_eq!(updated.n_cent, 5);
assert_eq!(updated.counts.len(), 5);
assert_eq!(updated.centroids.len(), 5 * base.dim as usize);
let bytes = crate::supertable::manifest::encoding::encode_cluster_centroids(&updated);
let decoded = crate::supertable::manifest::encoding::decode_cluster_centroids(&bytes)
.expect("split grid must reopen from wire bytes");
assert_eq!(decoded.n_cent, 5);
assert_eq!(decoded.centroids.len(), 5 * base.dim as usize);
}
#[test]
fn modality_primitives_separate_and_count_k() {
let dim = 64usize;
let threshold = 4.0;
let d_sample = |rows: &[EncodedCellRow], seed: u64| -> f64 {
let refs: Vec<&EncodedCellRow> = rows.iter().collect();
let decoded = decode_rows(&refs, dim);
let c = kmeans(&decoded, dim, 2, SPLIT_KMEANS_ITERS, seed);
ashman_d(&decoded, dim, &c)
};
let k_of = |rows: &[EncodedCellRow], seed: u64| -> usize {
let refs: Vec<&EncodedCellRow> = rows.iter().collect();
let decoded = decode_rows(&refs, dim);
let idx: Vec<usize> = (0..rows.len()).collect();
recursive_binary_k(&decoded, dim, &idx, seed, threshold, MODALITY_MAX_DEPTH)
};
for (sigma, seed) in [(0.1f32, 11u64), (0.03, 13)] {
let uni = synth_gaussian_cell(dim, 1, 1000, sigma, seed);
let d = d_sample(&uni, 0);
assert!(
(2.0..4.0).contains(&d),
"unimodal (sigma {sigma}) D should sit near the ~3 baseline, got {d}"
);
assert_eq!(
k_of(&uni, 0),
1,
"unimodal cell -> k = 1, got {}",
k_of(&uni, 0)
);
}
let bi = synth_gaussian_cell(dim, 2, 700, 0.02, 12);
assert!(
d_sample(&bi, 0) > 100.0,
"separated modes should score far above the baseline, got {}",
d_sample(&bi, 0)
);
let tri = synth_gaussian_cell(dim, 3, 700, 0.02, 21);
assert_eq!(
k_of(&tri, 0),
3,
"three separated modes -> k = 3, got {}",
k_of(&tri, 0)
);
}
#[test]
fn plan_sq8_split_separates_two_blobs() {
let dim = 4usize;
let mut rows = synth_rows(dim, 10, 0.0);
rows.extend(synth_rows(dim, 10, 10.0));
let clusters = synth_centroids(4, dim as u32);
let refs: Vec<&EncodedCellRow> = rows.iter().collect();
let (c0, c1, assign) = plan_sq8_split(&refs, &clusters, 1, Metric::L2Sq);
assert_eq!(c0.len(), dim);
assert_eq!(c1.len(), dim);
let dist: f32 = (0..dim).map(|d| (c0[d] - c1[d]).abs()).sum();
assert!(dist > 1.0, "split centroids should separate, got {dist}");
assert_eq!(assign.len(), rows.len());
assert_ne!(
assign[0],
assign[rows.len() - 1],
"the two separated blobs should split across sub-cells"
);
}
#[test]
fn plan_fixed_residual_split_preserves_payloads() {
let dim = 4usize;
let mut rows = synth_fixed_rows(dim, 10, 64);
rows.extend(synth_fixed_rows(dim, 10, 192));
let before: Vec<(Vec<u8>, Vec<u8>)> = rows
.iter()
.map(|row| (row.codes.clone(), row.residuals.clone()))
.collect();
let clusters = synth_centroids(4, dim as u32);
let refs: Vec<&EncodedCellRow> = rows.iter().collect();
let (left, right, _assign) = plan_sq8_split(&refs, &clusters, 1, Metric::Cosine);
let separation: f32 = left.iter().zip(&right).map(|(a, b)| (a - b).abs()).sum();
assert!(separation > 1.0);
let after: Vec<(Vec<u8>, Vec<u8>)> = rows
.iter()
.map(|row| (row.codes.clone(), row.residuals.clone()))
.collect();
assert_eq!(after, before);
}
fn assert_kway_split_balanced(
dim: usize,
n_blobs: usize,
per_blob: usize,
k: usize,
seed: u64,
) {
let rows = synth_gaussian_cell(dim, n_blobs, per_blob, 0.05, seed);
let n = rows.len();
let clusters = synth_centroids(1, dim as u32);
let refs: Vec<&EncodedCellRow> = rows.iter().collect();
let (cents, assign) = plan_sq8_split_kway(&refs, &clusters, 0, Metric::L2Sq, k, true);
let kk = (cents.len() / dim).max(1);
let mut counts = vec![0usize; kk];
for &a in &assign {
counts[a as usize] += 1;
}
let mean = n / kk;
let max_child = *counts.iter().max().expect("kk >= 1");
let empty = counts.iter().filter(|&&c| c == 0).count();
let route_faithful = assign
.iter()
.enumerate()
.filter(|&(i, &a)| {
let rv = dequantize_row(refs[i], dim);
let nearest = (0..kk)
.min_by(|&x, &y| {
distance(Metric::L2Sq, &rv, ¢s[x * dim..(x + 1) * dim])
.partial_cmp(&distance(
Metric::L2Sq,
&rv,
¢s[y * dim..(y + 1) * dim],
))
.unwrap_or(Ordering::Equal)
})
.unwrap_or(0);
a as usize == nearest
})
.count();
let route_frac = route_faithful as f64 / n as f64;
let mut sorted = counts.clone();
sorted.sort_unstable_by(|a, b| b.cmp(a));
eprintln!(
"[split-test] blobs={n_blobs} n={n} k_req={k} k_used={kk} mean={mean} max={max_child} \
empty={empty} route_fidelity={route_frac:.3} cells(desc)={sorted:?}",
);
assert_eq!(
empty, 0,
"no empty sub-cells; blobs={n_blobs} kk={kk} got {counts:?}"
);
assert!(
max_child <= n.div_ceil(k) + 1,
"over cap_target (blobs={n_blobs} k_req={k} k_used={kk}): max {max_child} vs {} \
got {counts:?}",
n.div_ceil(k),
);
assert!(
route_frac >= 0.95,
"low route-fidelity {route_frac:.3} (blobs={n_blobs} k_req={k} k_used={kk}) — \
self-tuning must raise k until most rows are in their nearest child"
);
}
#[test]
fn plan_sq8_split_kway_kmeans_balances_many_equal_blobs() {
let dim = 1024usize; assert_kway_split_balanced(dim, 16, 100, 10, 42);
assert_kway_split_balanced(dim, 32, 100, 20, 7);
assert_kway_split_balanced(dim, 64, 60, 40, 101);
}
#[test]
fn plan_sq8_split_kway_self_tunes_k_upward() {
let dim = 1024usize;
let rows = synth_gaussian_cell(dim, 16, 100, 0.05, 42);
let clusters = synth_centroids(1, dim as u32);
let refs: Vec<&EncodedCellRow> = rows.iter().collect();
let (cents, assign) = plan_sq8_split_kway(&refs, &clusters, 0, Metric::L2Sq, 2, true);
let kk = cents.len() / dim;
let populated = {
let mut seen = vec![false; kk.max(1)];
for &a in &assign {
seen[a as usize] = true;
}
seen.iter().filter(|&&s| s).count()
};
eprintln!("[self-tune] k_req=2 k_used={kk} populated={populated}");
assert!(
kk > 2,
"self-tuning must raise k above 2 on a 16-group cell, got {kk}"
);
assert!(
populated >= 2,
"at least 2 sub-cells populated, got {populated}"
);
}
}