Skip to main content

semantic_check/
semantic_check.rs

1// Semantic sanity check for a BERT-style embedding model.
2// Embeds a few sentences and prints the cosine-similarity matrix.
3// Related pairs should score clearly higher than unrelated pairs.
4//   cargo run --release -p rlx-embed --example semantic_check -- <model_dir>
5use rlx_embed::{BertTokenizer, Pooling, RlxBertModel, embed_with_rlx};
6use std::path::Path;
7
8fn cosine(a: &[f32], b: &[f32]) -> f32 {
9    let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
10    let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
11    let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
12    dot / (na * nb + 1e-8)
13}
14
15fn main() -> anyhow::Result<()> {
16    let dir = std::env::args()
17        .nth(1)
18        .expect("usage: semantic_check <model_dir>");
19    let dir = Path::new(&dir);
20    let pooling = if dir.to_string_lossy().to_lowercase().contains("bge") {
21        Pooling::Cls
22    } else {
23        Pooling::Mean
24    };
25
26    let tok = BertTokenizer::from_dir(dir, 64)?;
27    let mut model = RlxBertModel::load(
28        &dir.join("config.json"),
29        dir.join("model.safetensors").to_str().unwrap(),
30    )?;
31
32    let texts = [
33        "A man is playing a guitar.",                  // 0
34        "Someone is strumming an acoustic guitar.",    // 1  ~ 0
35        "It is freezing outside and snow is falling.", // 2
36        "The weather is cold and it is snowing.",      // 3  ~ 2
37    ];
38    let vecs = embed_with_rlx(&mut model, &tok, &texts, pooling)?;
39    println!("dim={} pooling={:?}", vecs[0].len(), pooling);
40
41    println!("\ncosine matrix:");
42    print!("      ");
43    for j in 0..texts.len() {
44        print!("  s{j:<5}");
45    }
46    println!();
47    for i in 0..texts.len() {
48        print!("s{i}  ");
49        for j in 0..texts.len() {
50            print!("  {:.3} ", cosine(&vecs[i], &vecs[j]));
51        }
52        println!();
53    }
54
55    let sim_related = (cosine(&vecs[0], &vecs[1]) + cosine(&vecs[2], &vecs[3])) / 2.0;
56    let sim_unrelated = (cosine(&vecs[0], &vecs[2])
57        + cosine(&vecs[0], &vecs[3])
58        + cosine(&vecs[1], &vecs[2])
59        + cosine(&vecs[1], &vecs[3]))
60        / 4.0;
61    println!("\nmean related  (0-1, 2-3): {sim_related:.3}");
62    println!("mean unrelated (cross)  : {sim_unrelated:.3}");
63    if sim_related > sim_unrelated + 0.05 {
64        println!("PASS: related pairs are clearly more similar than unrelated pairs.");
65    } else {
66        println!("FAIL: semantic ordering not observed.");
67        std::process::exit(1);
68    }
69    Ok(())
70}