#[path = "arcface/mod.rs"]
mod common;
use coremlit::embeddings::face::{AlignedFace, FaceEmbedder, FaceEmbedderOptions, arcface};
const SOLVE_TOLERANCE: f64 = 1e-9;
#[test]
fn the_reference_corpus_loads_and_its_vectors_are_raw() {
let reference = common::load_reference();
assert_eq!(reference.faces.len(), 18, "18 committed crops");
assert_eq!(reference.dim, 512);
assert_eq!(
reference
.faces
.iter()
.map(|f| f.person.as_str())
.collect::<std::collections::BTreeSet<_>>()
.len(),
6,
"six identities"
);
for face in &reference.faces {
let measured = face
.reference
.iter()
.map(|v| f64::from(*v) * f64::from(*v))
.sum::<f64>()
.sqrt();
assert!(
(measured - face.reference_norm).abs() < 1e-3,
"{}: recorded norm {} but the vector measures {measured}",
face.id,
face.reference_norm
);
assert!(
measured > 2.0,
"{}: the reference must be RAW; norm {measured:.4}",
face.id
);
}
}
#[test]
fn the_reference_directions_are_distinct_across_identities() {
let reference = common::load_reference();
let mut worst = -1.0f64;
for (i, a) in reference.faces.iter().enumerate() {
for b in reference.faces.iter().skip(i + 1) {
if a.person == b.person {
continue;
}
let cos = common::cosine(&a.reference, &b.reference);
assert!(
cos < common::SANITY_COS,
"{} and {} embed to the same direction (cos {cos:.6}); the parity gate could not tell \
them apart",
a.id,
b.id
);
worst = worst.max(cos);
}
}
eprintln!("[arcface] reference: closest different-person pair = {worst:.6}");
}
#[test]
fn the_rust_solve_reproduces_every_committed_alignment_matrix() {
let reference = common::load_reference();
let mut worst = 0.0f64;
for face in &reference.faces {
let solved = common::solved_transform(face).matrix();
for (index, (got, want)) in solved.iter().zip(face.align_matrix).enumerate() {
let allowed = SOLVE_TOLERANCE * want.abs().max(1.0);
let delta = (got - want).abs();
worst = worst.max(delta / want.abs().max(1.0));
assert!(
delta <= allowed,
"{}: matrix[{index}] solved {got} against the oracle's {want} (delta {delta:.3e}, \
allowed {allowed:.3e})",
face.id
);
}
}
eprintln!("[arcface] solve: worst relative deviation from the oracle = {worst:.3e}");
}
#[test]
#[ignore = "requires the staged arcface model (FACEKIT_TEST_MODELS)"]
fn the_door_reproduces_the_onnx_reference_on_the_recommended_arm() {
let reference = common::load_reference();
let embedder = FaceEmbedder::load(
common::model_path(),
arcface::MODEL,
FaceEmbedderOptions::new().with_compute(arcface::RECOMMENDED_COMPUTE),
)
.expect("load embedder");
let faces: Vec<AlignedFace> = reference
.faces
.iter()
.map(|face| {
AlignedFace::from_template_pixels(&face.crop)
.unwrap_or_else(|e| panic!("{}: wrap crop: {e}", face.id))
})
.collect();
let embeddings = embedder.embed(&faces).expect("embed the whole corpus");
assert_eq!(embeddings.len(), reference.faces.len());
let mut worst = (1.0f64, String::new());
for (face, embedding) in reference.faces.iter().zip(&embeddings) {
assert_eq!(embedding.dim(), reference.dim, "{}: width", face.id);
assert!(
embedding.as_slice().iter().all(|v| v.is_finite()),
"{}: non-finite component",
face.id
);
let cos = common::cosine(embedding.as_slice(), &face.reference);
eprintln!(
"[arcface] {:?} {}: cos vs ONNX fp32 = {cos:.8}",
arcface::RECOMMENDED_COMPUTE,
face.id
);
if cos < worst.0 {
worst = (cos, face.id.clone());
}
}
eprintln!(
"[arcface] {:?} worst cos = {:.8} ({}), 1-cos = {:.3e}, floor {}",
arcface::RECOMMENDED_COMPUTE,
worst.0,
worst.1,
1.0 - worst.0,
common::SANITY_COS
);
assert!(
worst.0 >= common::SANITY_COS,
"{}: cos {:.8} < {} on {:?}",
worst.1,
worst.0,
common::SANITY_COS,
arcface::RECOMMENDED_COMPUTE
);
}
#[test]
#[ignore = "requires the staged arcface model (FACEKIT_TEST_MODELS)"]
fn cross_face_geometry_matches_the_reference() {
const MAX_PAIR_DELTA: f64 = 1e-2;
let reference = common::load_reference();
let embedder = FaceEmbedder::load(
common::model_path(),
arcface::MODEL,
FaceEmbedderOptions::new().with_compute(arcface::RECOMMENDED_COMPUTE),
)
.expect("load embedder");
let faces: Vec<AlignedFace> = reference
.faces
.iter()
.map(|face| AlignedFace::from_template_pixels(&face.crop).expect("wrap crop"))
.collect();
let ours = embedder.embed(&faces).expect("embed the whole corpus");
let mut worst = (0.0f64, String::new());
for (i, a) in reference.faces.iter().enumerate() {
for (j, b) in reference.faces.iter().enumerate().skip(i + 1) {
let theirs = common::cosine(&a.reference, &b.reference);
let mine = common::cosine(ours[i].as_slice(), ours[j].as_slice());
let delta = (mine - theirs).abs();
if delta > worst.0 {
worst = (delta, format!("{}/{}", a.id, b.id));
}
assert!(
delta <= MAX_PAIR_DELTA,
"{}/{}: cross-face cosine moved by {delta:.3e} (reference {theirs:.8}, door {mine:.8})",
a.id,
b.id
);
}
}
eprintln!(
"[arcface] worst pairwise-cosine drift vs the reference: {:.3e} ({})",
worst.0, worst.1
);
}