use crate::embedding::colbert::{
DocIdAlignment, TokenEmbeddings, build_doc_token_ids, load_doc_id_alignment,
};
use crate::embedding::static_table::StaticTokenTable;
use anyhow::{Context, Result};
use ndarray::Array2;
use std::path::Path;
const WINDOW_LEN: usize = 5;
const CENTER_OFFSET: usize = 2;
pub struct StaticTokenEmbedder {
table: StaticTokenTable,
alignment: DocIdAlignment,
}
impl StaticTokenEmbedder {
pub fn new(model_dir: &Path) -> Result<Self> {
let table_path = crate::embedding::model_manager::static_token_table_path(model_dir);
let table = StaticTokenTable::load(&table_path).with_context(|| {
format!("failed to load static token table {}", table_path.display())
})?;
let alignment = load_doc_id_alignment(model_dir)?;
Ok(Self { table, alignment })
}
pub fn encode_documents(&self, texts: &[String]) -> Result<Vec<TokenEmbeddings>> {
texts
.iter()
.map(|text| {
let ids = build_doc_token_ids(&self.alignment, text)?;
Ok(mix_document(&self.table, &ids))
})
.collect()
}
}
pub(crate) fn window_ids_at<const N: usize>(ids: &[u32], i: usize) -> [u32; N] {
let center = N / 2;
let id = ids[i];
let mut window = [id; N];
for offset in 1..=center {
if i >= offset {
window[center - offset] = ids[i - offset];
}
if i + offset < ids.len() {
window[center + offset] = ids[i + offset];
}
}
window
}
fn normalize_in_place(v: &mut [f32]) {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in v.iter_mut() {
*x /= norm;
}
}
}
fn mix_token(
table: &StaticTokenTable,
window: [u32; WINDOW_LEN],
weights: [f32; WINDOW_LEN],
) -> Vec<f32> {
let dims = table.dims;
let rows: [Option<&[f32]>; WINDOW_LEN] = std::array::from_fn(|k| table.lookup(window[k]));
if rows[CENTER_OFFSET].is_some() {
let mut v = vec![0.0f32; dims];
for (k, row) in rows.iter().enumerate() {
if let Some(row) = row {
let w = weights[k];
for (vi, &ri) in v.iter_mut().zip(row.iter()) {
*vi += w * ri;
}
}
}
normalize_in_place(&mut v);
v
} else {
let mut sum = vec![0.0f32; dims];
let mut hits: u32 = 0;
for (k, row) in rows.iter().enumerate() {
if k == CENTER_OFFSET {
continue;
}
if let Some(row) = row {
hits += 1;
for (si, &ri) in sum.iter_mut().zip(row.iter()) {
*si += ri;
}
}
}
if hits == 0 {
return vec![0.0f32; dims]; }
let inv = 1.0 / hits as f32;
for s in &mut sum {
*s *= inv;
}
normalize_in_place(&mut sum);
sum
}
}
fn mix_document(table: &StaticTokenTable, ids: &[u32]) -> TokenEmbeddings {
let dims = table.dims;
let mut out = Array2::<f32>::zeros((ids.len(), dims));
for i in 0..ids.len() {
let window = window_ids_at::<WINDOW_LEN>(ids, i);
let row = mix_token(table, window, table.mix_weights);
out.row_mut(i).assign(&ndarray::ArrayView1::from(&row));
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn hand_built_table() -> StaticTokenTable {
let mut t = StaticTokenTable::new(6, 3, [0.1, 0.2, 0.4, 0.2, 0.1]);
t.set_row(1, &[1.0, 0.0, 0.0]);
t.set_row(2, &[0.0, 1.0, 0.0]);
t.set_row(3, &[0.0, 0.0, 1.0]);
t.set_row(4, &[1.0, 1.0, 0.0]);
t.set_row(5, &[0.0, 1.0, 1.0]);
t
}
fn normalize(v: &[f32]) -> Vec<f32> {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
v.iter().map(|x| x / norm).collect()
} else {
v.to_vec()
}
}
fn assert_close(got: &[f32], want: &[f32], msg: &str) {
assert_eq!(got.len(), want.len(), "{msg}: length mismatch");
for (g, w) in got.iter().zip(want.iter()) {
assert!((g - w).abs() < 1e-5, "{msg}: got {got:?}, want {want:?}");
}
}
#[test]
fn interior_position_matches_hand_computed_weighted_mix() {
let table = hand_built_table();
let ids = vec![1u32, 2, 3, 4, 5];
let doc = mix_document(&table, &ids);
let w = table.mix_weights;
let rows: Vec<&[f32]> = vec![
table.lookup(1).unwrap(),
table.lookup(2).unwrap(),
table.lookup(3).unwrap(),
table.lookup(4).unwrap(),
table.lookup(5).unwrap(),
];
let mut expected = vec![0.0f32; 3];
for (k, row) in rows.iter().enumerate() {
for (e, &r) in expected.iter_mut().zip(row.iter()) {
*e += w[k] * r;
}
}
let expected = normalize(&expected);
assert_close(
doc.row(2).as_slice().unwrap(),
&expected,
"interior mixed row",
);
}
#[test]
fn rows_are_l2_normalized() {
let table = hand_built_table();
let ids = vec![1u32, 2, 3, 4, 5];
let doc = mix_document(&table, &ids);
for i in 0..ids.len() {
let row = doc.row(i);
let norm: f32 = row.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-5,
"row {i} norm {norm} is not ~1.0: {row:?}"
);
}
}
#[test]
fn edge_position_reuses_center_id_per_task3_convention() {
let table = hand_built_table();
let ids = vec![1u32, 2, 3, 4, 5];
let doc = mix_document(&table, &ids);
let w = table.mix_weights;
let row1 = table.lookup(1).unwrap(); let row2 = table.lookup(2).unwrap(); let row3 = table.lookup(3).unwrap(); let mut expected = vec![0.0f32; 3];
for (i, &v) in row1.iter().enumerate() {
expected[i] += (w[0] + w[1] + w[2]) * v;
}
for (i, &v) in row2.iter().enumerate() {
expected[i] += w[3] * v;
}
for (i, &v) in row3.iter().enumerate() {
expected[i] += w[4] * v;
}
let expected = normalize(&expected);
assert_close(
doc.row(0).as_slice().unwrap(),
&expected,
"edge-position mixed row (i=0)",
);
}
#[test]
fn oov_center_falls_back_to_unmixed_mean_of_neighbor_hits() {
let table = hand_built_table();
let ids = vec![1u32, 2, 0, 4, 5];
let doc = mix_document(&table, &ids);
let neighbors = [
table.lookup(1).unwrap(),
table.lookup(2).unwrap(),
table.lookup(4).unwrap(),
table.lookup(5).unwrap(),
];
let mut mean = vec![0.0f32; 3];
for row in &neighbors {
for (m, &v) in mean.iter_mut().zip(row.iter()) {
*m += v;
}
}
for m in &mut mean {
*m /= neighbors.len() as f32;
}
let expected = normalize(&mean);
assert_close(
doc.row(2).as_slice().unwrap(),
&expected,
"OOV-center fallback row",
);
}
#[test]
fn everything_missing_emits_all_zero_row() {
let table = hand_built_table();
let ids = vec![0u32];
let doc = mix_document(&table, &ids);
let row = doc.row(0);
assert!(
row.iter().all(|&v| v == 0.0),
"expected an all-zero row when every window lookup misses, got {row:?}"
);
}
fn expect_err(result: Result<StaticTokenEmbedder>) -> anyhow::Error {
match result {
Ok(_) => panic!("expected an error, got Ok"),
Err(e) => e,
}
}
#[test]
fn new_errors_when_static_token_table_missing() {
let tmp = tempfile::TempDir::new().unwrap();
let err = expect_err(StaticTokenEmbedder::new(tmp.path()));
assert!(
err.to_string().contains("static_token_table.bin"),
"expected a static-token-table error, got: {err}"
);
}
#[test]
fn new_errors_when_tokenizer_missing() {
let tmp = tempfile::TempDir::new().unwrap();
let table = StaticTokenTable::new(4, 2, [0.0, 0.0, 1.0, 0.0, 0.0]);
table
.save(&tmp.path().join("static_token_table.bin"))
.unwrap();
let err = expect_err(StaticTokenEmbedder::new(tmp.path()));
assert!(
err.to_string().to_lowercase().contains("tokenizer"),
"expected a tokenizer-loading error, got: {err}"
);
}
fn test_tokenizer_dir() -> Option<std::path::PathBuf> {
let dir = crate::config::SemantexConfig::default()
.models_dir()
.join("LateOn-Code-edge");
(dir.join("tokenizer.json").exists() && dir.join("onnx_config.json").exists())
.then_some(dir)
}
#[test]
fn end_to_end_with_real_tokenizer_and_hand_built_table() {
let Some(model_dir) = test_tokenizer_dir() else {
return;
};
let vocab_size = crate::embedding::colbert::ColbertEmbedder::new(&model_dir)
.unwrap()
.tokenizer_vocab_size()
.unwrap();
let mut table = StaticTokenTable::new(vocab_size, 4, [0.1, 0.2, 0.4, 0.2, 0.1]);
table.set_row(0, &[1.0, 0.0, 0.0, 0.0]);
let tmp = tempfile::TempDir::new().unwrap();
table
.save(&tmp.path().join("static_token_table.bin"))
.unwrap();
std::fs::copy(
model_dir.join("tokenizer.json"),
tmp.path().join("tokenizer.json"),
)
.unwrap();
std::fs::copy(
model_dir.join("onnx_config.json"),
tmp.path().join("onnx_config.json"),
)
.unwrap();
let embedder = StaticTokenEmbedder::new(tmp.path()).unwrap();
let texts = vec!["fn main() { println!(\"hi\"); }".to_string()];
let out = embedder.encode_documents(&texts).unwrap();
assert_eq!(out.len(), 1);
assert!(
out[0].nrows() > 0,
"document should produce at least one token row"
);
assert_eq!(out[0].ncols(), 4, "row width must match the table's dims");
}
}