semantic_check/
semantic_check.rs1use 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.", "Someone is strumming an acoustic guitar.", "It is freezing outside and snow is falling.", "The weather is cold and it is snowing.", ];
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}