use std::sync::LazyLock;
use serde::Deserialize;
pub const REFERENCE_UNRELATED: f32 = 0.4128;
pub const REFERENCE_PARAPHRASE: f32 = 0.8176;
const REFERENCE_SPAN: f32 = REFERENCE_PARAPHRASE - REFERENCE_UNRELATED;
const MIN_THRESHOLD: f32 = 0.0;
const MAX_THRESHOLD: f32 = 1.0;
const PROBES_JSON: &str = include_str!("embed_probes.json");
#[derive(Debug, Default, Deserialize)]
struct ProbeCorpus {
paraphrase: Vec<(String, String)>,
unrelated: Vec<(String, String)>,
}
impl ProbeCorpus {
fn pairs(&self) -> impl Iterator<Item = &(String, String)> {
self.paraphrase.iter().chain(self.unrelated.iter())
}
}
static CORPUS: LazyLock<ProbeCorpus> =
LazyLock::new(|| serde_json::from_str(PROBES_JSON).unwrap_or_default());
pub fn probe_texts() -> Vec<String> {
CORPUS
.pairs()
.flat_map(|(a, b)| [a.clone(), b.clone()])
.collect()
}
pub fn measure(vectors: &[Vec<f32>]) -> Option<Calibration> {
let (paraphrase_pairs, unrelated_pairs) = (CORPUS.paraphrase.len(), CORPUS.unrelated.len());
if paraphrase_pairs == 0 || unrelated_pairs == 0 {
return None;
}
if vectors.len() != 2 * (paraphrase_pairs + unrelated_pairs) {
return None;
}
let dim = vectors[0].len();
if dim == 0 || vectors.iter().any(|v| v.len() != dim) {
return None;
}
let mean = |pairs: &[Vec<f32>]| -> Option<f32> {
let sum: f32 = pairs.chunks_exact(2).map(|p| cosine(&p[0], &p[1])).sum();
let mean = sum / (pairs.len() / 2) as f32;
mean.is_finite().then_some(mean)
};
let split = 2 * paraphrase_pairs;
Some(Calibration {
paraphrase: mean(&vectors[..split])?,
unrelated: mean(&vectors[split..])?,
})
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Calibration {
pub unrelated: f32,
pub paraphrase: f32,
}
#[derive(Debug, Clone, Copy)]
pub struct SimilarityScale {
affine: Option<Affine>,
}
#[derive(Debug, Clone, Copy)]
struct Affine {
unrelated: f32,
factor: f32,
}
impl SimilarityScale {
pub fn identity() -> Self {
Self { affine: None }
}
pub fn from_calibration(c: Calibration) -> Self {
let plausible = [c.unrelated, c.paraphrase]
.iter()
.all(|v| v.is_finite() && (-1.0..=1.0).contains(v));
let span = c.paraphrase - c.unrelated;
if !plausible || span <= 0.0 {
return Self::identity();
}
Self {
affine: Some(Affine {
unrelated: c.unrelated,
factor: span / REFERENCE_SPAN,
}),
}
}
pub fn map(&self, threshold: f32) -> f32 {
match self.affine {
None => threshold,
Some(Affine { unrelated, factor }) => (unrelated
+ (threshold - REFERENCE_UNRELATED) * factor)
.clamp(MIN_THRESHOLD, MAX_THRESHOLD),
}
}
}
fn cosine(a: &[f32], b: &[f32]) -> f32 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if na == 0.0 || nb == 0.0 {
return 0.0;
}
dot / (na * nb)
}
#[cfg(test)]
mod tests {
use super::*;
const GATES: [f32; 3] = [0.85, 0.72, 0.62];
const E5_UNRELATED: f32 = 0.7897;
const E5_PARAPHRASE: f32 = 0.9456;
fn scale(unrelated: f32, paraphrase: f32) -> SimilarityScale {
SimilarityScale::from_calibration(Calibration {
unrelated,
paraphrase,
})
}
#[test]
fn corpus_is_well_formed() {
assert_eq!(CORPUS.paraphrase.len(), 8, "8 paraphrase pairs (§8.2)");
assert_eq!(CORPUS.unrelated.len(), 8, "8 unrelated pairs (§8.2)");
assert!(
CORPUS.pairs().all(|(a, b)| !a.is_empty() && !b.is_empty()),
"an empty probe would embed to nothing and skew a mean"
);
}
#[test]
fn probe_texts_are_flattened_in_a_stable_order() {
let first = probe_texts();
assert_eq!(first.len(), 32, "16 pairs, both halves");
assert_eq!(
first,
probe_texts(),
"the order must not vary between calls"
);
assert_eq!(first[0], CORPUS.paraphrase[0].0);
assert_eq!(first[1], CORPUS.paraphrase[0].1);
assert_eq!(first[16], CORPUS.unrelated[0].0);
}
fn probe_vectors() -> Vec<Vec<f32>> {
let mut v = Vec::new();
for _ in 0..CORPUS.paraphrase.len() {
v.push(vec![1.0, 0.0]);
v.push(vec![1.0, 1.0]);
}
for _ in 0..CORPUS.unrelated.len() {
v.push(vec![1.0, 0.0]);
v.push(vec![0.0, 1.0]);
}
v
}
#[test]
fn measure_returns_the_two_means() {
let c = measure(&probe_vectors()).expect("a well-formed batch");
assert!(
(c.paraphrase - std::f32::consts::FRAC_1_SQRT_2).abs() < 1e-5,
"{c:?}"
);
assert!(c.unrelated.abs() < 1e-5, "{c:?}");
}
#[test]
fn measure_refuses_a_batch_that_does_not_describe_the_corpus() {
assert_eq!(measure(&[]), None, "empty");
assert_eq!(measure(&vec![vec![1.0, 0.0]; 31]), None, "one short");
assert_eq!(measure(&vec![vec![1.0, 0.0]; 33]), None, "one long");
let mut empty_vector = probe_vectors();
empty_vector[3] = Vec::new();
assert_eq!(measure(&empty_vector), None, "an empty vector");
let mut mixed = probe_vectors();
mixed[3] = vec![1.0, 0.0, 0.0];
assert_eq!(measure(&mixed), None, "mixed dimensionalities");
}
#[test]
fn identity_passes_thresholds_through_untouched() {
for t in GATES.iter().chain(&[0.0, 0.5, 1.0]) {
assert_eq!(SimilarityScale::identity().map(*t), *t);
}
}
#[test]
fn the_reference_calibration_is_the_identity() {
let s = scale(REFERENCE_UNRELATED, REFERENCE_PARAPHRASE);
for t in GATES {
assert!((s.map(t) - t).abs() < 1e-5, "{t} -> {}", s.map(t));
}
}
#[test]
fn e5_reproduces_the_measured_thresholds() {
let s = scale(E5_UNRELATED, E5_PARAPHRASE);
for (raw, expected) in [(0.85, 0.958), (0.72, 0.908), (0.62, 0.869)] {
assert!(
(s.map(raw) - expected).abs() < 0.002,
"{raw} -> {} (§8.2 says {expected})",
s.map(raw)
);
}
let trait_gate = s.map(0.72);
assert!(
E5_UNRELATED > 0.72 && trait_gate > E5_UNRELATED,
"unrelated {E5_UNRELATED} must sit above the raw gate and below the mapped one ({trait_gate})"
);
}
#[test]
fn a_narrower_range_moves_every_gate_up_and_keeps_their_order() {
let s = scale(E5_UNRELATED, E5_PARAPHRASE);
let mapped: Vec<f32> = GATES.iter().map(|t| s.map(*t)).collect();
assert!(
mapped.windows(2).all(|w| w[0] > w[1]),
"the gates keep their relative order: {mapped:?}"
);
assert!(
GATES.iter().zip(&mapped).all(|(raw, m)| m > raw),
"a compressed range pushes every gate up: {mapped:?}"
);
}
#[test]
fn a_degenerate_calibration_falls_back_to_the_identity() {
let degenerate = [
(0.9, 0.4, "inverted: paraphrase below unrelated"),
(
0.5,
0.5,
"zero span: every gate would collapse onto one value",
),
(f32::NAN, 0.8, "not a number"),
(0.4, f32::NAN, "not a number"),
(0.4, f32::INFINITY, "not finite"),
(-2.0, 0.8, "outside the cosine range"),
(0.4, 1.5, "outside the cosine range"),
];
for (unrelated, paraphrase, why) in degenerate {
for t in GATES {
assert_eq!(scale(unrelated, paraphrase).map(t), t, "{why}");
}
}
}
#[test]
fn a_mapped_threshold_never_leaves_the_cosine_range() {
let s = scale(-1.0, 1.0);
for t in GATES.iter().chain(&[0.0, 0.1, 0.99, 1.0]) {
let mapped = s.map(*t);
assert!(
(MIN_THRESHOLD..=MAX_THRESHOLD).contains(&mapped),
"{t} -> {mapped}"
);
}
}
#[test]
fn cosine_is_scale_invariant_and_safe_on_bad_input() {
assert!((cosine(&[1.0, 2.0], &[2.0, 4.0]) - 1.0).abs() < 1e-6);
assert_eq!(cosine(&[1.0, 0.0], &[0.0, 1.0]), 0.0);
assert_eq!(cosine(&[1.0], &[1.0, 0.0]), 0.0, "length mismatch");
assert_eq!(cosine(&[], &[]), 0.0, "empty");
assert_eq!(cosine(&[0.0, 0.0], &[1.0, 0.0]), 0.0, "zero norm");
}
}