use crate::embedding::colbert::{ColbertEmbedder, TokenEmbeddings};
use crate::embedding::static_table::StaticTokenTable;
use anyhow::Result;
const WINDOW_LEN: usize = 5;
const CENTER_OFFSET: usize = 2;
const RESERVOIR_CAPACITY: usize = 200_000;
const RESERVOIR_SEED: u64 = 0xE813_9B0C_5A70_C13B;
const SINGULAR_EPS: f64 = 1e-9;
const FALLBACK_MIX_WEIGHTS: [f32; WINDOW_LEN] = [0.0, 0.0, 1.0, 0.0, 0.0];
pub trait DocTokenEncoder {
fn encode_with_ids(&self, texts: &[String]) -> Result<Vec<(Vec<u32>, TokenEmbeddings)>>;
fn vocab_size(&self) -> usize;
}
impl DocTokenEncoder for ColbertEmbedder {
fn encode_with_ids(&self, texts: &[String]) -> Result<Vec<(Vec<u32>, TokenEmbeddings)>> {
self.encode_documents_with_ids(texts)
}
fn vocab_size(&self) -> usize {
self.tokenizer_vocab_size().unwrap_or_else(|e| {
panic!("ColbertEmbedder::vocab_size: failed to load tokenizer/config: {e}")
})
}
}
struct ReservoirItem {
window_ids: [u32; WINDOW_LEN],
center_embedding: Vec<f32>,
}
struct SplitMix64(u64);
impl SplitMix64 {
fn new(seed: u64) -> Self {
Self(seed)
}
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn next_below(&mut self, bound: u64) -> u64 {
self.next_u64() % bound
}
}
struct Accumulator {
vocab_size: usize,
dims: Option<usize>,
sums: Option<Vec<f32>>,
counts: Vec<u64>,
reservoir: Vec<ReservoirItem>,
seen: u64,
rng: SplitMix64,
}
impl Accumulator {
fn new(vocab_size: usize) -> Self {
Self {
vocab_size,
dims: None,
sums: None,
counts: vec![0u64; vocab_size],
reservoir: Vec::new(),
seen: 0,
rng: SplitMix64::new(RESERVOIR_SEED),
}
}
fn ingest_batch(&mut self, encoder: &dyn DocTokenEncoder, texts: &[String]) -> Result<()> {
for (ids, emb) in encoder.encode_with_ids(texts)? {
self.ingest_document(&ids, &emb)?;
}
Ok(())
}
fn ingest_document(&mut self, ids: &[u32], emb: &TokenEmbeddings) -> Result<()> {
if ids.is_empty() {
return Ok(());
}
anyhow::ensure!(
ids.len() == emb.nrows(),
"distill: token id/embedding row mismatch ({} ids vs {} rows); \
DocTokenEncoder contract violated",
ids.len(),
emb.nrows()
);
let ncols = emb.ncols();
match self.dims {
None => {
self.dims = Some(ncols);
self.sums = Some(vec![0.0f32; self.vocab_size * ncols]);
}
Some(d) => anyhow::ensure!(
d == ncols,
"distill: inconsistent embedding width across batches ({d} vs {ncols}); \
DocTokenEncoder must report a constant width"
),
}
let dims = self.dims.expect("just set above");
for i in 0..ids.len() {
let id = ids[i];
let idx = id as usize;
anyhow::ensure!(
idx < self.vocab_size,
"distill: token id {id} out of bounds for vocab_size {}",
self.vocab_size
);
let row = emb.row(i);
{
let sums = self.sums.as_mut().expect("just set above");
let base = idx * dims;
for (s, v) in sums[base..base + dims].iter_mut().zip(row.iter()) {
*s += v;
}
}
self.counts[idx] += 1;
let mut window_ids = [id; WINDOW_LEN];
for offset in 1..=CENTER_OFFSET {
if i >= offset {
window_ids[CENTER_OFFSET - offset] = ids[i - offset];
}
if i + offset < ids.len() {
window_ids[CENTER_OFFSET + offset] = ids[i + offset];
}
}
self.reservoir_sample(ReservoirItem {
window_ids,
center_embedding: row.to_vec(),
});
}
Ok(())
}
fn reservoir_sample(&mut self, item: ReservoirItem) {
if self.reservoir.len() < RESERVOIR_CAPACITY {
self.reservoir.push(item);
} else {
let j = self.rng.next_below(self.seen + 1);
if (j as usize) < RESERVOIR_CAPACITY {
self.reservoir[j as usize] = item;
}
}
self.seen += 1;
}
}
fn solve_5x5(
mut a: [[f64; WINDOW_LEN]; WINDOW_LEN],
mut b: [f64; WINDOW_LEN],
) -> Option<[f32; WINDOW_LEN]> {
for col in 0..WINDOW_LEN {
let mut pivot_row = col;
let mut pivot_val = a[col][col].abs();
for (row, arow) in a.iter().enumerate().skip(col + 1) {
if arow[col].abs() > pivot_val {
pivot_val = arow[col].abs();
pivot_row = row;
}
}
if pivot_val < SINGULAR_EPS {
return None;
}
if pivot_row != col {
a.swap(col, pivot_row);
b.swap(col, pivot_row);
}
let pivot_row_vals = a[col];
let pivot_b = b[col];
for (row, arow) in a.iter_mut().enumerate().skip(col + 1) {
let factor = arow[col] / pivot_row_vals[col];
if factor == 0.0 {
continue;
}
for (k, &pv) in pivot_row_vals.iter().enumerate().skip(col) {
arow[k] -= factor * pv;
}
b[row] -= factor * pivot_b;
}
}
let mut x = [0.0f64; WINDOW_LEN];
for row in (0..WINDOW_LEN).rev() {
let mut sum = b[row];
for (k, &xk) in x.iter().enumerate().skip(row + 1) {
sum -= a[row][k] * xk;
}
x[row] = sum / a[row][row];
}
Some(std::array::from_fn(|i| x[i] as f32))
}
fn fit_mix_weights(reservoir: &[ReservoirItem], mean_table: &[Vec<f32>]) -> [f32; WINDOW_LEN] {
let mut a = [[0.0f64; WINDOW_LEN]; WINDOW_LEN];
let mut b = [0.0f64; WINDOW_LEN];
let dot = |u: &[f32], v: &[f32]| -> f64 {
u.iter()
.zip(v.iter())
.map(|(&a, &b)| f64::from(a) * f64::from(b))
.sum()
};
for item in reservoir {
let xs: [&[f32]; WINDOW_LEN] =
std::array::from_fn(|k| mean_table[item.window_ids[k] as usize].as_slice());
for k in 0..WINDOW_LEN {
b[k] += dot(xs[k], &item.center_embedding);
for l in k..WINDOW_LEN {
let akl = dot(xs[k], xs[l]);
a[k][l] += akl;
if l != k {
a[l][k] += akl;
}
}
}
}
solve_5x5(a, b).unwrap_or(FALLBACK_MIX_WEIGHTS)
}
pub fn distill(
encoder: &dyn DocTokenEncoder,
corpus: impl Iterator<Item = String>,
batch: usize,
) -> Result<StaticTokenTable> {
anyhow::ensure!(batch > 0, "distill: `batch` must be greater than zero");
let vocab_size = encoder.vocab_size();
anyhow::ensure!(vocab_size > 0, "distill: encoder reports vocab_size = 0");
let mut acc = Accumulator::new(vocab_size);
let mut batch_buf: Vec<String> = Vec::with_capacity(batch);
for text in corpus {
batch_buf.push(text);
if batch_buf.len() == batch {
acc.ingest_batch(encoder, &batch_buf)?;
batch_buf.clear();
}
}
if !batch_buf.is_empty() {
acc.ingest_batch(encoder, &batch_buf)?;
}
let dims = acc
.dims
.ok_or_else(|| anyhow::anyhow!("distill: empty corpus produced no tokens"))?;
let sums = acc.sums.expect("sums is set alongside dims");
let mut mean_table: Vec<Vec<f32>> = vec![vec![0.0; dims]; vocab_size];
for (token_id, (&count, mean_row)) in acc.counts.iter().zip(mean_table.iter_mut()).enumerate() {
if count == 0 {
continue;
}
let base = token_id * dims;
let inv = 1.0 / count as f32;
let mut row: Vec<f32> = sums[base..base + dims].iter().map(|&v| v * inv).collect();
let norm = row.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > 0.0 {
for v in &mut row {
*v /= norm;
}
}
*mean_row = row;
}
let mix_weights = fit_mix_weights(&acc.reservoir, &mean_table);
let mut table = StaticTokenTable::new(vocab_size, dims, mix_weights);
for (token_id, (&count, mean_row)) in acc.counts.iter().zip(mean_table.iter()).enumerate() {
if count > 0 {
table.set_row(token_id as u32, mean_row);
}
}
Ok(table)
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::Array2;
struct FakeEncoder {
dims: usize,
vocab_size: usize,
}
impl FakeEncoder {
const LEAK: f32 = 0.05;
fn new(vocab_size: usize, dims: usize) -> Self {
Self { dims, vocab_size }
}
fn embed_row(&self, token_id: u32, prev_id: u32) -> Vec<f32> {
let mut row = vec![0.0f32; self.dims];
row[token_id as usize % self.dims] += 1.0;
let leak_pos = (token_id as usize + 1 + prev_id as usize) % self.dims;
row[leak_pos] += Self::LEAK;
row
}
}
impl DocTokenEncoder for FakeEncoder {
fn encode_with_ids(&self, texts: &[String]) -> Result<Vec<(Vec<u32>, TokenEmbeddings)>> {
Ok(texts
.iter()
.map(|text| {
let bytes = text.as_bytes();
let mut ids: Vec<u32> = Vec::with_capacity(bytes.len());
let mut emb = Array2::<f32>::zeros((bytes.len(), self.dims));
for (i, &b) in bytes.iter().enumerate() {
let id = u32::from(b - b'a');
let prev = if i == 0 { id } else { ids[i - 1] };
let row = self.embed_row(id, prev);
for (d, v) in row.into_iter().enumerate() {
emb[[i, d]] = v;
}
ids.push(id);
}
(ids, emb)
})
.collect())
}
fn vocab_size(&self) -> usize {
self.vocab_size
}
}
#[test]
fn table_rows_are_the_normalized_mean_of_emitted_embeddings() {
let dims = 4;
let vocab = 4; let encoder = FakeEncoder::new(vocab, dims);
let corpus = vec!["abcd".to_string(), "dcba".to_string(), "aabb".to_string()];
let mut sums = vec![vec![0.0f32; dims]; vocab];
let mut counts = vec![0u64; vocab];
for text in &corpus {
let bytes = text.as_bytes();
let mut ids: Vec<u32> = Vec::new();
for (i, &b) in bytes.iter().enumerate() {
let id = u32::from(b - b'a');
let prev = if i == 0 { id } else { ids[i - 1] };
let row = encoder.embed_row(id, prev);
for (s, v) in sums[id as usize].iter_mut().zip(row.iter()) {
*s += v;
}
counts[id as usize] += 1;
ids.push(id);
}
}
let table = distill(&encoder, corpus.into_iter(), 2).expect("non-empty corpus distills");
for token_id in 0..vocab as u32 {
let count = counts[token_id as usize];
assert!(count > 0, "test corpus must exercise every token id");
let mean: Vec<f32> = sums[token_id as usize]
.iter()
.map(|&s| s / count as f32)
.collect();
let norm = mean.iter().map(|v| v * v).sum::<f32>().sqrt();
let expected: Vec<f32> = mean.iter().map(|&v| v / norm).collect();
let row = table
.lookup(token_id)
.unwrap_or_else(|| panic!("token {token_id} missing from table"));
for (got, want) in row.iter().zip(expected.iter()) {
assert!(
(got - want).abs() < 1e-4,
"token {token_id}: got {row:?}, want {expected:?}"
);
}
}
}
#[test]
fn table_rows_are_l2_normalized() {
let dims = 4;
let vocab = 4;
let encoder = FakeEncoder::new(vocab, dims);
let corpus = vec!["abcd".to_string(), "dcba".to_string(), "aabb".to_string()];
let table = distill(&encoder, corpus.into_iter(), 2).expect("non-empty corpus distills");
for token_id in 0..vocab as u32 {
let row = table
.lookup(token_id)
.unwrap_or_else(|| panic!("token {token_id} missing from table"));
let norm = row.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-4,
"token {token_id} row norm {norm} is not ~1.0 ({row:?})"
);
}
}
#[test]
fn center_mix_weight_dominates_when_context_effect_is_small() {
let dims = 6;
let vocab = 5; let encoder = FakeEncoder::new(vocab, dims);
let base = "abcdeabcdeabcdeabcdeabcde";
let corpus: Vec<String> = (0..50)
.map(|i| {
let start = i % base.len();
format!("{}{}", &base[start..], &base[..start])
})
.collect();
let table = distill(&encoder, corpus.into_iter(), 8).expect("non-empty corpus distills");
let w = table.mix_weights;
assert_ne!(
w, FALLBACK_MIX_WEIGHTS,
"weights exactly match the singular-system fallback; the reservoir fit \
likely didn't run (this corpus should produce a well-conditioned system)"
);
for (k, &wk) in w.iter().enumerate() {
if k != 2 {
assert!(
w[2] > wk,
"center weight w[2]={} should dominate w[{k}]={wk}; full weights: {w:?}",
w[2]
);
}
}
}
#[test]
fn singular_reservoir_falls_back_to_pure_center_lookup() {
let encoder = FakeEncoder::new(4, 4);
let corpus: Vec<String> = vec!["a".repeat(20); 5];
let table = distill(&encoder, corpus.into_iter(), 8).expect("non-empty corpus distills");
assert_eq!(
table.mix_weights, FALLBACK_MIX_WEIGHTS,
"a single-token corpus produces a rank-1 normal matrix and must fall back exactly"
);
}
#[test]
fn empty_corpus_errors_instead_of_producing_a_zero_table() {
let encoder = FakeEncoder::new(4, 4);
let corpus: Vec<String> = Vec::new();
let err = distill(&encoder, corpus.into_iter(), 8).unwrap_err();
assert!(
err.to_string().to_lowercase().contains("empty"),
"expected an 'empty corpus' error, got: {err}"
);
}
#[test]
fn corpus_of_only_empty_documents_errors_like_an_empty_corpus() {
let encoder = FakeEncoder::new(4, 4);
let corpus = vec![String::new(), String::new()];
let err = distill(&encoder, corpus.into_iter(), 8).unwrap_err();
assert!(
err.to_string().to_lowercase().contains("empty"),
"expected an 'empty corpus' error, got: {err}"
);
}
#[test]
fn mismatched_ids_and_embedding_rows_errors() {
struct BrokenEncoder;
impl DocTokenEncoder for BrokenEncoder {
fn encode_with_ids(
&self,
texts: &[String],
) -> Result<Vec<(Vec<u32>, TokenEmbeddings)>> {
Ok(texts
.iter()
.map(|_| (vec![0, 1, 2], Array2::<f32>::zeros((2, 4))))
.collect())
}
fn vocab_size(&self) -> usize {
4
}
}
let err = distill(&BrokenEncoder, vec!["x".to_string()].into_iter(), 8).unwrap_err();
assert!(
err.to_string().contains("mismatch"),
"expected a row/id mismatch error, got: {err}"
);
}
#[test]
fn out_of_bounds_token_id_errors() {
struct BrokenEncoder;
impl DocTokenEncoder for BrokenEncoder {
fn encode_with_ids(
&self,
texts: &[String],
) -> Result<Vec<(Vec<u32>, TokenEmbeddings)>> {
Ok(texts
.iter()
.map(|_| (vec![99], Array2::<f32>::zeros((1, 4))))
.collect())
}
fn vocab_size(&self) -> usize {
4
}
}
let err = distill(&BrokenEncoder, vec!["x".to_string()].into_iter(), 8).unwrap_err();
assert!(
err.to_string().contains("out of bounds"),
"expected an out-of-bounds token id error, got: {err}"
);
}
}