use std::{
cmp::Ordering,
collections::HashMap,
sync::{
Mutex, PoisonError,
atomic::{AtomicU32, Ordering as AtomicOrdering},
},
};
use crate::{
config,
superfile::vector::{
cell_posting::{EncodedCellRow, MaterializedIvfRow, manifest_centroid_components_from_row},
distance::{
Metric, distance, nearest_k_centroids_bytes, nearest_k_centroids_transposed, normalize,
relative_score_window,
},
kmeans::{kmeans, kmeans_pp},
quant::BitQuantizer,
reader::CellFineCalibrationView,
reservoir::Reservoir,
rotation::RandomRotation,
spill::SpilledCellRows,
},
supertable::{
error::BuildError,
manifest::{
ClusterCentroids, RABITQ_ADMIT_CELL_SHORTLIST_FRACTION,
RABITQ_ADMIT_CELL_SHORTLIST_MIN, RabitqAdmitContext,
list::{WIDTH_LAW_KS, WIDTH_LAW_MAX_K},
},
},
};
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 row_fp = dequantize_row(row, clusters.dim as usize);
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_row_into(row, &mut out);
out
}
fn dequantize_row_into(row: &EncodedCellRow, out: &mut [f32]) {
let dim = out.len();
row.rerank_codec
.ops()
.expect("encoded row uses a quantized-rerank codec")
.dequantize_row_into(
&row.codes,
&row.residuals,
dim,
&row.scale,
&row.offset,
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)
}
pub(crate) const WIDTH_LAW_QUERY_SAMPLE: usize = 256;
const ACCEPTANCE_BAR_RECALL: f64 = 0.99;
const WIDTH_LAW_CONFIDENCE_Z: f64 = 0.0;
pub(crate) const RERANK_LAW_POOL_CELLS: usize = 64;
const RERANK_LAW_POOL_MARGIN: usize = 2;
pub(crate) fn rerank_pool_hint(width_for_k: &[u32; WIDTH_LAW_KS.len()], n_cent: usize) -> usize {
let widest = width_for_k.iter().copied().max().unwrap_or(0) as usize;
(widest * RERANK_LAW_POOL_MARGIN)
.max(RERANK_LAW_POOL_CELLS)
.min(n_cent.max(1))
}
const RERANK_LAW_EST_BINS: usize = 4096;
const WIDTH_LAW_SAMPLE_SEED: u64 = 0x51ED_CA1B;
const WIDTH_LAW_SCORE_CHUNK: usize = 1024;
struct WidthLawQueries {
queries: Vec<f32>,
ids: Vec<i128>,
}
pub(crate) struct WidthLawCalibration {
dim: usize,
metric: Metric,
reservoir: Reservoir,
slot_ids: Vec<i128>,
dequant_scratch: Vec<f32>,
frozen: Option<WidthLawQueries>,
tops: Mutex<Vec<Vec<(f32, u32, i128, f32)>>>,
fine_ranks: Mutex<HashMap<(u32, i128, u32), u32>>,
max_fine: AtomicU32,
pool_cells: usize,
target_recall: f64,
rerank: Option<RerankLawObservation>,
}
struct RerankLawObservation {
quant: BitQuantizer,
q_rot: Vec<f32>,
q_total: Vec<f32>,
q_l1: Vec<f32>,
pools: Vec<Vec<u32>>,
hist: Mutex<Vec<Vec<u64>>>,
}
impl RerankLawObservation {
fn bin(&self, qi: usize, est: f32) -> usize {
let l1 = self.q_l1[qi];
if l1 <= 0.0 {
return RERANK_LAW_EST_BINS - 1;
}
let frac = ((l1 - est) / (2.0 * l1)).clamp(0.0, 1.0);
((frac * RERANK_LAW_EST_BINS as f32) as usize).min(RERANK_LAW_EST_BINS - 1)
}
}
pub(crate) struct CalibratedLaws {
pub(crate) width_for_k: [u32; WIDTH_LAW_KS.len()],
pub(crate) fine_for_k: [u32; WIDTH_LAW_KS.len()],
pub(crate) rerank_for_k: [u32; WIDTH_LAW_KS.len()],
pub(crate) pool_cells: u32,
}
#[cfg(test)]
mod pool_hint_tests {
use super::*;
#[test]
fn rerank_pool_hint_scales_with_the_stamped_width() {
assert_eq!(rerank_pool_hint(&[0, 0, 0, 0], 256), RERANK_LAW_POOL_CELLS);
assert_eq!(rerank_pool_hint(&[1, 2, 8, 16], 256), RERANK_LAW_POOL_CELLS);
assert_eq!(rerank_pool_hint(&[33, 79, 97, 104], 256), 208);
assert_eq!(rerank_pool_hint(&[33, 79, 97, 104], 150), 150);
assert_eq!(rerank_pool_hint(&[1, 0, 0, 0], 0), 1);
}
}
pub(crate) fn clear_rerank_beyond_pool(
width_for_k: &[u32; WIDTH_LAW_KS.len()],
rerank_for_k: &mut [u32; WIDTH_LAW_KS.len()],
pool_cells: &[u32; WIDTH_LAW_KS.len()],
) {
for ((w, r), pool) in width_for_k
.iter()
.zip(rerank_for_k.iter_mut())
.zip(pool_cells.iter())
{
if *w > *pool {
*r = 0;
}
}
}
pub(crate) fn merge_rerank_with_pools(
rerank: &mut [u32; WIDTH_LAW_KS.len()],
pools: &mut [u32; WIDTH_LAW_KS.len()],
measured: &[u32; WIDTH_LAW_KS.len()],
measured_pool: u32,
) {
for ((slot, pool), m) in rerank.iter_mut().zip(pools.iter_mut()).zip(measured.iter()) {
if *m > *slot {
*slot = *m;
*pool = measured_pool;
} else if *m == *slot && *m > 0 {
*pool = (*pool).max(measured_pool);
}
}
}
fn floor_monotone(law: &mut [u32; WIDTH_LAW_KS.len()]) {
let mut floor = 0u32;
for w in law.iter_mut().filter(|w| **w > 0) {
*w = (*w).max(floor);
floor = *w;
}
}
impl WidthLawCalibration {
pub(crate) fn new(dim: usize, metric: Metric, target_recall: f64) -> Self {
let valid = |v: f64| v.is_finite() && v > 0.0 && v <= 1.0;
let fallback_target = {
let configured = config::global().vector.target_recall;
if valid(configured) {
configured
} else {
ACCEPTANCE_BAR_RECALL
}
};
let target_recall = if valid(target_recall) {
target_recall
} else {
tracing::warn!(
target_recall,
fallback = fallback_target,
"vector.target_recall must be in (0, 1]; falling back to the default"
);
fallback_target
};
Self {
dim,
metric,
reservoir: Reservoir::new(WIDTH_LAW_QUERY_SAMPLE, dim, WIDTH_LAW_SAMPLE_SEED),
slot_ids: Vec::with_capacity(WIDTH_LAW_QUERY_SAMPLE),
dequant_scratch: vec![0f32; dim],
frozen: None,
tops: Mutex::new(Vec::new()),
fine_ranks: Mutex::new(HashMap::new()),
max_fine: AtomicU32::new(0),
pool_cells: RERANK_LAW_POOL_CELLS,
target_recall,
rerank: None,
}
}
pub(crate) fn offer(&mut self, row: &MaterializedIvfRow) {
debug_assert!(self.frozen.is_none(), "offer after freeze");
dequantize_row_into(&row.encoded, &mut self.dequant_scratch);
if let Some(slot) = self.reservoir.update_traced(&self.dequant_scratch) {
if slot == self.slot_ids.len() {
self.slot_ids.push(row.stable_id);
} else {
self.slot_ids[slot] = row.stable_id;
}
}
}
pub(crate) fn freeze(&mut self, grid: &ClusterCentroids, rot_seed: u64, pool_cells: usize) {
self.pool_cells = pool_cells
.max(RERANK_LAW_POOL_CELLS)
.min((grid.n_cent as usize).max(1));
let queries = self.reservoir.sample().to_vec();
let ids = self.slot_ids.clone();
let n_queries = ids.len();
*self.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![Vec::new(); n_queries];
if n_queries > 0 && grid.n_cent > 0 {
let rotation = RandomRotation::new(self.dim, rot_seed);
let mut q_rot = vec![0f32; n_queries * self.dim];
let mut q_total = Vec::with_capacity(n_queries);
let mut q_l1 = Vec::with_capacity(n_queries);
let mut pools = Vec::with_capacity(n_queries);
for (qi, q) in queries.chunks_exact(self.dim).enumerate() {
let out = &mut q_rot[qi * self.dim..(qi + 1) * self.dim];
rotation.apply(q, out);
q_total.push(out.iter().sum());
q_l1.push(out.iter().map(|v| v.abs()).sum());
let mut pool: Vec<u32> = grid
.rank_cells(self.metric, q)
.into_iter()
.take(self.pool_cells)
.map(|(cell, _)| cell)
.collect();
pool.sort_unstable();
pools.push(pool);
}
self.rerank = Some(RerankLawObservation {
quant: BitQuantizer::new(self.dim),
q_rot,
q_total,
q_l1,
pools,
hist: Mutex::new(vec![Vec::new(); n_queries]),
});
}
self.frozen = Some(WidthLawQueries { queries, ids });
}
pub(crate) fn score_cell(&self, cell: u32, spill: &SpilledCellRows) -> Result<(), BuildError> {
let Some(frozen) = self.frozen.as_ref() else {
return Err(BuildError::Store(
"width-law score_cell before freeze".into(),
));
};
let n_queries = frozen.ids.len();
if n_queries == 0 {
return Ok(());
}
let mut partial: Vec<Vec<(f32, u32, i128, f32)>> = vec![Vec::new(); n_queries];
let members = self.pool_members(cell);
let mut hist_local: HashMap<usize, Vec<u64>> = HashMap::new();
let mut reader = spill.reader()?;
let mut remaining = spill.n_rows();
let mut scratch = vec![0f32; self.dim];
while remaining > 0 {
let chunk = reader.next_chunk(WIDTH_LAW_SCORE_CHUNK.min(remaining))?;
remaining -= chunk.len();
self.score_slice(
frozen,
cell,
&chunk,
&members,
&mut scratch,
&mut partial,
&mut hist_local,
);
}
self.merge_partial(partial, hist_local);
Ok(())
}
pub(crate) fn score_rows(
&self,
cell: u32,
rows: &[MaterializedIvfRow],
) -> Result<(), BuildError> {
let Some(frozen) = self.frozen.as_ref() else {
return Err(BuildError::Store(
"width-law score_rows before freeze".into(),
));
};
let n_queries = frozen.ids.len();
if n_queries == 0 {
return Ok(());
}
let mut partial: Vec<Vec<(f32, u32, i128, f32)>> = vec![Vec::new(); n_queries];
let members = self.pool_members(cell);
let mut hist_local: HashMap<usize, Vec<u64>> = HashMap::new();
let mut scratch = vec![0f32; self.dim];
for chunk in rows.chunks(WIDTH_LAW_SCORE_CHUNK) {
self.score_slice(
frozen,
cell,
chunk,
&members,
&mut scratch,
&mut partial,
&mut hist_local,
);
}
self.merge_partial(partial, hist_local);
Ok(())
}
fn merge_partial(
&self,
partial: Vec<Vec<(f32, u32, i128, f32)>>,
hist_local: HashMap<usize, Vec<u64>>,
) {
let k_max = WIDTH_LAW_MAX_K;
let mut tops = self.tops.lock().unwrap_or_else(PoisonError::into_inner);
for (qi, cand) in partial.into_iter().enumerate() {
merge_candidates(&mut tops[qi], cand, k_max);
}
drop(tops);
if let Some(rl) = self.rerank.as_ref()
&& !hist_local.is_empty()
{
let mut hist = rl.hist.lock().unwrap_or_else(PoisonError::into_inner);
for (qi, delta) in hist_local {
let slot = &mut hist[qi];
if slot.is_empty() {
*slot = delta;
} else {
for (a, b) in slot.iter_mut().zip(delta) {
*a = a.saturating_add(b);
}
}
}
}
}
fn pool_members(&self, cell: u32) -> Vec<usize> {
let Some(rl) = self.rerank.as_ref() else {
return Vec::new();
};
(0..rl.pools.len())
.filter(|&qi| rl.pools[qi].binary_search(&cell).is_ok())
.collect()
}
fn score_slice(
&self,
frozen: &WidthLawQueries,
cell: u32,
rows: &[MaterializedIvfRow],
members: &[usize],
scratch: &mut [f32],
partial: &mut [Vec<(f32, u32, i128, f32)>],
hist_local: &mut HashMap<usize, Vec<u64>>,
) {
let k_max = WIDTH_LAW_MAX_K;
let rl = self.rerank.as_ref();
let mut est_of = vec![f32::NEG_INFINITY; frozen.ids.len()];
for row in rows {
if let Some(rl) = rl
&& row.rabitq_code.len() == rl.quant.code_bytes()
{
for &qi in members {
if row.stable_id == frozen.ids[qi] {
continue;
}
let q_rot = &rl.q_rot[qi * self.dim..(qi + 1) * self.dim];
let est = rl.quant.estimate_dot_rotated_with_total(
q_rot,
&row.rabitq_code,
rl.q_total[qi],
);
if est.is_finite() {
est_of[qi] = est;
let bins = hist_local
.entry(qi)
.or_insert_with(|| vec![0u64; RERANK_LAW_EST_BINS]);
let bin = rl.bin(qi, est);
bins[bin] = bins[bin].saturating_add(1);
}
}
}
dequantize_row_into(&row.encoded, scratch);
if self.metric == Metric::Cosine {
normalize(scratch);
}
for (qi, q) in frozen.queries.chunks_exact(self.dim).enumerate() {
if row.stable_id == frozen.ids[qi] {
continue;
}
partial[qi].push((
distance(self.metric, q, scratch),
cell,
row.stable_id,
est_of[qi],
));
}
if rl.is_some() {
for &qi in members {
est_of[qi] = f32::NEG_INFINITY;
}
}
}
for cand in partial.iter_mut() {
truncate_ascending(cand, k_max);
}
}
pub(crate) fn observe_shard_views(&self, views: &[CellFineCalibrationView]) {
let Some(frozen) = self.frozen.as_ref() else {
return;
};
if frozen.ids.is_empty() {
return;
}
let per_cell: HashMap<u32, Vec<(u32, i128)>> = {
let tops = self.tops.lock().unwrap_or_else(PoisonError::into_inner);
let mut map: HashMap<u32, Vec<(u32, i128)>> = HashMap::new();
for (qi, cands) in tops.iter().enumerate() {
for &(_, cell, id, _) in cands {
map.entry(cell).or_default().push((qi as u32, id));
}
}
map
};
for view in views {
let Some(cell_id) = view.cell_id else {
continue;
};
let Some(cands) = per_cell.get(&cell_id) else {
continue;
};
if view.n_fine == 0 || view.dim != self.dim {
continue;
}
self.max_fine
.fetch_max(view.n_fine as u32, AtomicOrdering::Relaxed);
let mut rank_cache: HashMap<u32, Vec<u32>> = HashMap::new();
let mut ranks = self
.fine_ranks
.lock()
.unwrap_or_else(PoisonError::into_inner);
for &(qi, id) in cands {
let Some(&cluster) = view.cluster_of_stable.get(&id) else {
continue;
};
let rank_of = rank_cache.entry(qi).or_insert_with(|| {
let q = &frozen.queries[qi as usize * self.dim..(qi as usize + 1) * self.dim];
let ranked = nearest_k_centroids_bytes(
self.metric,
q,
&view.fine_centroids_bytes,
view.n_fine,
view.dim,
view.n_fine,
);
let mut rank_of = vec![0u32; view.n_fine];
for (rank, (c, _)) in ranked.iter().enumerate() {
rank_of[*c as usize] = rank as u32;
}
rank_of
});
if let Some(&r) = rank_of.get(cluster as usize) {
ranks.insert((qi, id, cell_id), r);
}
}
}
}
pub(crate) fn finish(self, grid: &ClusterCentroids) -> Option<CalibratedLaws> {
let frozen = self.frozen?;
let n_queries = frozen.ids.len();
if n_queries == 0 || grid.n_cent == 0 {
return None;
}
let n_cells = grid.n_cent as usize;
let tops = self
.tops
.into_inner()
.unwrap_or_else(PoisonError::into_inner);
let mut law = [0u32; WIDTH_LAW_KS.len()];
let mut coverage_sums: Vec<Vec<f64>> = vec![vec![0f64; n_cells]; WIDTH_LAW_KS.len()];
let mut coverage_sq_sums: Vec<Vec<f64>> = vec![vec![0f64; n_cells]; WIDTH_LAW_KS.len()];
let mut support = [0usize; WIDTH_LAW_KS.len()];
let fine_ranks = self
.fine_ranks
.lock()
.unwrap_or_else(PoisonError::into_inner);
let max_fine = self.max_fine.load(AtomicOrdering::Relaxed).max(1) as usize;
let mut fine_law = [0u32; WIDTH_LAW_KS.len()];
let mut fine_sums: Vec<Vec<f64>> = vec![vec![0f64; max_fine]; WIDTH_LAW_KS.len()];
let rerank_prefix: Option<Vec<Vec<u64>>> = self.rerank.as_ref().map(|rl| {
rl.hist
.lock()
.unwrap_or_else(PoisonError::into_inner)
.iter()
.map(|h| {
let mut run = 0u64;
h.iter()
.map(|&c| {
run = run.saturating_add(c);
run
})
.collect()
})
.collect()
});
let mut rerank_law = [0u32; WIDTH_LAW_KS.len()];
let mut rerank_ranks: [Vec<u64>; WIDTH_LAW_KS.len()] = Default::default();
let mut rank_of_cell = vec![0u32; n_cells];
for (qi, cand) in tops.iter().enumerate() {
let q = &frozen.queries[qi * self.dim..(qi + 1) * self.dim];
let ranked = grid.rank_cells(self.metric, q);
for (rank, (cell, _)) in ranked.iter().enumerate() {
if let Some(slot) = rank_of_cell.get_mut(*cell as usize) {
*slot = rank as u32;
}
}
let mut sorted = cand.clone();
sorted.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
for (ki, &k) in WIDTH_LAW_KS.iter().enumerate() {
if sorted.len() < k {
continue;
}
support[ki] += 1;
let mut per_rank = vec![0u32; n_cells];
for (_, cell, _, _) in &sorted[..k] {
let &rank = rank_of_cell.get(*cell as usize)?;
per_rank[rank as usize] += 1;
}
let mut covered = 0u32;
for (rank, count) in per_rank.iter().enumerate() {
covered += count;
let x = f64::from(covered) / k as f64;
coverage_sums[ki][rank] += x;
coverage_sq_sums[ki][rank] += x * x;
}
let mut per_fine_rank = vec![0u32; max_fine];
for (_, cell, id, _) in &sorted[..k] {
if let Some(&r) = fine_ranks.get(&(qi as u32, *id, *cell)) {
per_fine_rank[(r as usize).min(max_fine - 1)] += 1;
}
}
if let (Some(rl), Some(prefix)) = (self.rerank.as_ref(), rerank_prefix.as_ref()) {
for (_, _, _, est) in &sorted[..k] {
if est.is_finite() && !prefix[qi].is_empty() {
rerank_ranks[ki].push(prefix[qi][rl.bin(qi, *est)]);
} else {
rerank_ranks[ki].push(u64::MAX);
}
}
}
let mut fine_covered = 0u32;
for (rank, count) in per_fine_rank.iter().enumerate() {
fine_covered += count;
fine_sums[ki][rank] += f64::from(fine_covered) / k as f64;
}
}
}
for (ki, sums) in coverage_sums.iter().enumerate() {
if support[ki] == 0 {
continue;
}
let n = support[ki] as f64;
let target = self.target_recall * n;
if let Some(rank) = sums.iter().enumerate().position(|(rank, &s)| {
let mean = s / n;
let var = (coverage_sq_sums[ki][rank] / n - mean * mean).max(0.0);
let se = (var / n).sqrt();
(mean - WIDTH_LAW_CONFIDENCE_Z * se) * n >= target
}) {
law[ki] = (rank + 1) as u32;
}
let stage_target = self.target_recall * support[ki] as f64;
if let Some(rank) = fine_sums[ki].iter().position(|&s| s >= stage_target) {
fine_law[ki] = (rank + 1) as u32;
}
}
for (ki, &w) in law.iter().enumerate() {
let ranks = &mut rerank_ranks[ki];
if w == 0 || w as usize > self.pool_cells || ranks.is_empty() {
continue;
}
ranks.sort_unstable();
let needed = (self.target_recall * ranks.len() as f64).ceil() as usize;
if let Some(&crossing) = ranks.get(needed.saturating_sub(1).min(ranks.len() - 1))
&& crossing != u64::MAX
{
rerank_law[ki] = crossing.min(u64::from(u32::MAX)) as u32;
}
}
floor_monotone(&mut law);
floor_monotone(&mut fine_law);
floor_monotone(&mut rerank_law);
(law.iter().any(|&w| w > 0)).then_some(CalibratedLaws {
width_for_k: law,
fine_for_k: fine_law,
rerank_for_k: rerank_law,
pool_cells: self.pool_cells as u32,
})
}
}
fn merge_candidates(
acc: &mut Vec<(f32, u32, i128, f32)>,
mut cand: Vec<(f32, u32, i128, f32)>,
cap: usize,
) {
acc.append(&mut cand);
acc.sort_unstable_by(|a, b| {
a.2.cmp(&b.2)
.then_with(|| a.0.total_cmp(&b.0))
.then_with(|| b.3.is_finite().cmp(&a.3.is_finite()))
.then_with(|| a.1.cmp(&b.1))
});
acc.dedup_by_key(|c| c.2);
truncate_ascending(acc, cap);
}
fn truncate_ascending(cand: &mut Vec<(f32, u32, i128, f32)>, cap: usize) {
if cand.len() > cap {
cand.sort_unstable_by(|a, b| a.0.total_cmp(&b.0));
cand.truncate(cap);
}
}
#[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())
}
#[cfg(test)]
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)
}
pub(crate) fn insert_split_centroids_batch(
base: &ClusterCentroids,
splits: &[(u32, &[f32], usize)],
) -> (ClusterCentroids, Vec<Vec<u32>>) {
debug_assert!(
{
let mut parents: Vec<u32> = splits.iter().map(|(p, _, _)| *p).collect();
parents.sort_unstable();
parents.windows(2).all(|w| w[0] != w[1])
},
"batch splits must target distinct parent cells"
);
let dim = base.dim as usize;
let old_n = base.n_cent as usize;
let appended: usize = splits.iter().map(|(_, _, k)| k - 1).sum();
let new_n = old_n + appended;
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));
}
let mut ids_per_split = Vec::with_capacity(splits.len());
let mut next_id = old_n;
for &(cell_id, sub_centroids, k) in splits {
debug_assert_eq!(sub_centroids.len(), k * dim);
let p = cell_id as usize;
fp32[p * dim..(p + 1) * dim].copy_from_slice(&sub_centroids[..dim]);
let mut ids = vec![cell_id];
for j in 1..k {
fp32[next_id * dim..(next_id + 1) * dim]
.copy_from_slice(&sub_centroids[j * dim..(j + 1) * dim]);
ids.push(next_id as u32);
next_id += 1;
}
ids_per_split.push(ids);
}
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_per_split)
}
#[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 bytes::Bytes;
use super::*;
#[test]
fn every_stage_crosses_at_the_configured_target() {
const SHIPPED_WIDTH: f64 = 0.99;
assert_eq!(shipped_target_recall(), SHIPPED_WIDTH);
for target in [0.80, 0.90, 0.95, 0.99] {
let cal = WidthLawCalibration::new(8, Metric::Cosine, target);
assert_eq!(
cal.target_recall, target,
"the calibration carries the target unmodified"
);
}
}
#[test]
fn an_out_of_range_target_falls_back_to_the_configured_default() {
for bad in [f64::NAN, 0.0, -0.5, 1.5, f64::INFINITY, f64::NEG_INFINITY] {
let cal = WidthLawCalibration::new(8, Metric::Cosine, bad);
assert_eq!(
cal.target_recall,
shipped_target_recall(),
"target {bad} must fall back, not be stored"
);
}
let cal = WidthLawCalibration::new(8, Metric::Cosine, 1.0);
assert_eq!(cal.target_recall, 1.0);
}
fn shipped_target_recall() -> f64 {
crate::config::global().vector.target_recall
}
fn fp32_le_bytes(vals: &[f32]) -> Bytes {
Bytes::from(
vals.iter()
.flat_map(|v| v.to_le_bytes())
.collect::<Vec<u8>>(),
)
}
#[test]
fn width_law_finish_dedups_replicated_rows() {
const DIM: usize = 4;
let grid = ClusterCentroids::from_fp32(
4,
DIM as u32,
&[
1.0, 0.0, 0.0, 0.0, 0.9, 0.1, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0,
],
vec![1; 4],
);
let mut cal = WidthLawCalibration::new(DIM, Metric::Cosine, shipped_target_recall());
let mut query = vec![0.0f32; DIM];
query[0] = 1.0;
cal.frozen = Some(WidthLawQueries {
queries: query,
ids: vec![999],
});
let ninf = f32::NEG_INFINITY;
let mut cands = vec![(0.01, 0, 1, ninf), (0.02, 1, 1, ninf), (0.03, 1, 1, ninf)];
cands
.extend((2..=8).map(|id| (0.03 + id as f32 * 0.01, (id % 2) as u32, id as i128, ninf)));
cands.push((0.5, 3, 9, ninf));
cands.push((0.6, 3, 10, ninf));
let mut acc = Vec::new();
merge_candidates(&mut acc, cands, WIDTH_LAW_MAX_K);
*cal.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![acc];
let law = cal
.finish(&grid)
.expect("law from planted candidates")
.width_for_k;
assert_eq!(law[0], 1, "top-1 coverage is the nearest cell");
assert_eq!(
law[1], 4,
"replicated copies must not pad top-k coverage (got width {})",
law[1]
);
assert_eq!(&law[2..], &[0, 0], "unsupported points stay uncalibrated");
}
#[test]
fn depth_law_walks_fine_ranks() {
const DIM: usize = 4;
let grid = ClusterCentroids::from_fp32(1, DIM as u32, &[1.0, 0.0, 0.0, 0.0], vec![1; 1]);
let mut cal = WidthLawCalibration::new(DIM, Metric::Cosine, shipped_target_recall());
let mut query = vec![0.0f32; DIM];
query[0] = 1.0;
cal.frozen = Some(WidthLawQueries {
queries: query,
ids: vec![999],
});
let cands: Vec<(f32, u32, i128, f32)> = (1..=10)
.map(|id| (id as f32 * 0.01, 0u32, id as i128, f32::NEG_INFINITY))
.collect();
let mut acc = Vec::new();
merge_candidates(&mut acc, cands, WIDTH_LAW_MAX_K);
*cal.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![acc];
let mut cluster_of_stable = HashMap::new();
cluster_of_stable.insert(1i128, 1u32);
for id in 2..=9i128 {
cluster_of_stable.insert(id, 0u32);
}
cluster_of_stable.insert(10i128, 2u32);
let view = CellFineCalibrationView {
cell_id: Some(0),
dim: DIM,
n_fine: 3,
fine_centroids_bytes: fp32_le_bytes(&[
1.0, 0.0, 0.0, 0.0, 0.6, 0.8, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0,
]),
cluster_of_stable,
};
cal.observe_shard_views(&[view]);
let laws = cal.finish(&grid).expect("laws from planted candidates");
assert_eq!(laws.width_for_k[..2], [1, 1], "one cell holds everything");
assert_eq!(
laws.fine_for_k[0], 2,
"top-1 sits in the rank-1 fine cluster"
);
assert_eq!(
laws.fine_for_k[1], 3,
"top-10 coverage needs the rank-2 cluster"
);
assert_eq!(&laws.fine_for_k[2..], &[0, 0], "unsupported points stay 0");
}
#[test]
fn rerank_points_clear_when_stamped_width_outgrows_pool() {
let width = [
1,
RERANK_LAW_POOL_CELLS as u32,
(RERANK_LAW_POOL_CELLS + 1) as u32,
500,
];
let mut rerank = [10, 20, 30, 40];
clear_rerank_beyond_pool(
&width,
&mut rerank,
&[RERANK_LAW_POOL_CELLS as u32; WIDTH_LAW_KS.len()],
);
assert_eq!(
rerank,
[10, 20, 0, 0],
"points at widths beyond the pool must fall back to the constant"
);
}
#[test]
fn rerank_law_reads_survivor_budget_from_histograms() {
const DIM: usize = 4;
let grid = ClusterCentroids::from_fp32(1, DIM as u32, &[1.0, 0.0, 0.0, 0.0], vec![1; 1]);
let mut cal = WidthLawCalibration::new(DIM, Metric::Cosine, shipped_target_recall());
let mut query = vec![0.0f32; DIM];
query[0] = 1.0;
cal.frozen = Some(WidthLawQueries {
queries: query,
ids: vec![999],
});
let rl = RerankLawObservation {
quant: BitQuantizer::new(DIM),
q_rot: vec![0.0; DIM],
q_total: vec![0.0],
q_l1: vec![1.0],
pools: vec![vec![0]],
hist: Mutex::new(vec![Vec::new()]),
};
let mut hist = vec![0u64; RERANK_LAW_EST_BINS];
hist[rl.bin(0, 0.9)] = 3;
hist[rl.bin(0, 0.5)] = 50;
*rl.hist.lock().unwrap_or_else(PoisonError::into_inner) = vec![hist];
cal.rerank = Some(rl);
let cands: Vec<(f32, u32, i128, f32)> = (1..=10)
.map(|id| {
let est = if id == 1 { 0.9 } else { 0.5 };
(id as f32 * 0.01, 0u32, id as i128, est)
})
.collect();
let mut acc = Vec::new();
merge_candidates(&mut acc, cands, WIDTH_LAW_MAX_K);
*cal.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![acc];
let laws = cal.finish(&grid).expect("laws from planted candidates");
assert_eq!(
laws.rerank_for_k[0], 3,
"k=1 budget = the best candidate's distractor count"
);
assert_eq!(
laws.rerank_for_k[1], 53,
"k=10 budget = the worst top-10 candidate's distractor count"
);
assert_eq!(
&laws.rerank_for_k[2..],
&[0, 0],
"unsupported points stay uncalibrated"
);
}
#[test]
fn depth_law_missing_rank_is_conservative() {
const DIM: usize = 4;
let grid = ClusterCentroids::from_fp32(1, DIM as u32, &[1.0, 0.0, 0.0, 0.0], vec![1; 1]);
let mut cal = WidthLawCalibration::new(DIM, Metric::Cosine, shipped_target_recall());
let mut query = vec![0.0f32; DIM];
query[0] = 1.0;
cal.frozen = Some(WidthLawQueries {
queries: query,
ids: vec![999],
});
let cands: Vec<(f32, u32, i128, f32)> = (1..=10)
.map(|id| (id as f32 * 0.01, 0u32, id as i128, f32::NEG_INFINITY))
.collect();
let mut acc = Vec::new();
merge_candidates(&mut acc, cands, WIDTH_LAW_MAX_K);
*cal.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![acc];
let mut cluster_of_stable = HashMap::new();
for id in 1..=9i128 {
cluster_of_stable.insert(id, 0u32);
}
let view = CellFineCalibrationView {
cell_id: Some(0),
dim: DIM,
n_fine: 2,
fine_centroids_bytes: fp32_le_bytes(&[
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0,
]),
cluster_of_stable,
};
cal.observe_shard_views(&[view]);
let laws = cal.finish(&grid).expect("laws from planted candidates");
assert_eq!(laws.fine_for_k[0], 1, "top-1 was observed at rank 0");
assert_eq!(
laws.fine_for_k[1], 0,
"k=10 misses id 10's rank: 9/10 < 0.99 coverage, point stays 0"
);
}
#[test]
fn width_law_supports_points_per_query_and_stays_monotone() {
const DIM: usize = 4;
let grid = ClusterCentroids::from_fp32(
4,
DIM as u32,
&[
1.0, 0.0, 0.0, 0.0, 0.9, 0.1, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0,
],
vec![1; 4],
);
let mut cal = WidthLawCalibration::new(DIM, Metric::Cosine, shipped_target_recall());
let mut queries = vec![0.0f32; 2 * DIM];
queries[0] = 1.0;
queries[DIM] = 1.0;
cal.frozen = Some(WidthLawQueries {
queries,
ids: vec![998, 999],
});
let k_max = WIDTH_LAW_MAX_K;
let a: Vec<(f32, u32, i128, f32)> = (1..=10)
.map(|id| {
let cell = match id {
1 | 2 => 0u32,
3 | 4 => 1,
5..=7 => 2,
_ => 3,
};
(id as f32 * 0.01, cell, id as i128, f32::NEG_INFINITY)
})
.collect();
let b: Vec<(f32, u32, i128, f32)> = (100..200)
.map(|id| (id as f32 * 0.001, 0u32, id as i128, f32::NEG_INFINITY))
.collect();
let (mut acc_a, mut acc_b) = (Vec::new(), Vec::new());
merge_candidates(&mut acc_a, a, k_max);
merge_candidates(&mut acc_b, b, k_max);
*cal.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![acc_a, acc_b];
let law = cal
.finish(&grid)
.expect("law from planted candidates")
.width_for_k;
assert_eq!(law[0], 1, "top-1: both queries covered by the nearest cell");
assert_eq!(
law[1], 4,
"k=10 measured over both queries needs A's spread"
);
assert_eq!(
law[2], 4,
"k=100: B alone measures width 1; the monotone floor lifts it \
to the k=10 width instead of stamping a recall inversion"
);
assert_eq!(law[3], 0, "k=1000 has no supporting query and stays 0");
}
#[test]
fn width_law_bails_on_out_of_range_cell_id() {
const DIM: usize = 4;
let grid = ClusterCentroids::from_fp32(
2,
DIM as u32,
&[
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0,
],
vec![1; 2],
);
let mut cal = WidthLawCalibration::new(DIM, Metric::Cosine, shipped_target_recall());
let mut query = vec![0.0f32; DIM];
query[0] = 1.0;
cal.frozen = Some(WidthLawQueries {
queries: query,
ids: vec![999],
});
let mut acc = Vec::new();
merge_candidates(
&mut acc,
vec![(0.01, 7, 1, f32::NEG_INFINITY)],
WIDTH_LAW_MAX_K,
);
*cal.tops.lock().unwrap_or_else(PoisonError::into_inner) = vec![acc];
assert!(
cal.finish(&grid).is_none(),
"inconsistent calibration input must abandon the law, not panic"
);
}
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 insert_split_centroids_batch_matches_sequential_fold() {
let dim = 8usize;
let base = synth_centroids(6, dim as u32);
let sub = |tag: f32, k: usize| -> Vec<f32> {
(0..k * dim).map(|i| tag + i as f32 * 0.01).collect()
};
let (s0, s2, s5) = (sub(1.0, 2), sub(2.0, 4), sub(3.0, 3));
let splits: Vec<(u32, &[f32], usize)> = vec![(0, &s0, 2), (2, &s2, 4), (5, &s5, 3)];
let (batched, batched_ids) = insert_split_centroids_batch(&base, &splits);
let mut folded = base.clone();
let mut folded_ids = Vec::new();
for &(cell, sub_centroids, k) in &splits {
let (next, ids) = insert_split_centroids(&folded, cell, sub_centroids, k);
folded = next;
folded_ids.push(ids);
}
assert_eq!(batched.n_cent, folded.n_cent);
assert_eq!(batched.dim, folded.dim);
assert_eq!(batched.centroids, folded.centroids);
assert_eq!(batched.counts, folded.counts);
assert_eq!(batched_ids, folded_ids);
assert_eq!(batched_ids[0], vec![0, 6]);
assert_eq!(batched_ids[1], vec![2, 7, 8, 9]);
assert_eq!(batched_ids[2], vec![5, 10, 11]);
assert_eq!(batched.counts.len(), 12);
assert!(batched.counts[6..].iter().all(|&c| c == 0));
let bytes = crate::supertable::manifest::encoding::encode_cluster_centroids(&batched);
let decoded = crate::supertable::manifest::encoding::decode_cluster_centroids(&bytes)
.expect("batch-split grid must reopen from wire bytes");
assert_eq!(decoded.n_cent, 12);
assert_eq!(decoded.centroids.len(), 12 * dim);
}
#[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}"
);
}
}