#[path = "../../support/workspace_root.rs"]
#[allow(dead_code)]
mod workspace_root;
#[allow(unused_imports)]
pub use workspace_root::{checkout_parent, models_root, workspace_root};
use std::path::{Path, PathBuf};
use coremlit::embeddings::face::{
ARCFACE_TEMPLATE, LANDMARK_COUNT, Point, SimilarityTransform, TEMPLATE_BYTES,
};
#[allow(dead_code)]
pub const HF_REPO: &str = "FinDIT-Studio/facekit-coreml";
#[allow(dead_code)]
pub const HF_REVISION: &str = "70e212696bd3c472e28718e2e39c79467b97805e";
#[allow(dead_code)]
pub const SOURCE_PACK: &str = "buffalo_l.zip";
#[allow(dead_code)]
pub const SOURCE_PACK_SHA256: &str =
"80ffe37d8a5940d59a7384c201a2a38d4741f2f3c51eef46ebb28218a7b0ca2f";
#[allow(dead_code)]
pub const SOURCE_MEMBER: &str = "w600k_r50.onnx";
#[allow(dead_code)]
pub const SOURCE_MEMBER_SHA256: &str =
"4c06341c33c2ca1f86781dab0e829f88ad5b64be9fba56e56bc9ebdefc619e43";
#[allow(dead_code)]
pub const INSIGHTFACE_REVISION: &str = "ffa12d315041c0505b077c7ff057ca914bb8dc7e";
#[allow(dead_code)]
pub const BUNDLE_NAME: &str = "w600k_r50.mlmodelc";
#[allow(dead_code)]
pub const ARTIFACT_SHA256: &[(&str, &str)] = &[
(
"analytics/coremldata.bin",
"1320de26a121f36a6dde0a1faab329d69006560709c0671cc1254cfccf4cdb5f",
),
(
"coremldata.bin",
"f95b43443adb213f1f96136b42e5113c539ce897efa916b8982b82f03e92a38f",
),
(
"metadata.json",
"a0dd2ec43b8182a7184a4116df29c2e10afd454be1f239e4b5e1efb507fa22b8",
),
(
"model.mil",
"050f69f10f5687971fb8f808d9da53b01a8d512c7013346fbc22daa948e42d26",
),
(
"weights/weight.bin",
"aa08d7826a70f9bc237ea0532a5eec12cb83b8375148a1b0650f104cbb2ff492",
),
];
#[allow(dead_code)]
pub const SANITY_COS: f64 = 0.99;
#[allow(dead_code)]
pub fn models_dir() -> PathBuf {
std::env::var_os("FACEKIT_TEST_MODELS").map_or_else(
|| workspace_root::models_root().join("facekit"),
PathBuf::from,
)
}
#[allow(dead_code)]
pub fn model_path() -> PathBuf {
models_dir().join(BUNDLE_NAME)
}
#[allow(dead_code)]
pub fn fixtures_dir() -> PathBuf {
Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("face")
.join("fixtures")
}
#[allow(dead_code)]
pub fn sha256_hex(bytes: &[u8]) -> String {
use core::fmt::Write;
use sha2::{Digest, Sha256};
Sha256::digest(bytes)
.iter()
.fold(String::new(), |mut acc, b| {
let _ = write!(acc, "{b:02x}");
acc
})
}
#[allow(dead_code)]
pub fn sha256_file(path: &Path) -> String {
let bytes = std::fs::read(path).unwrap_or_else(|e| panic!("read {path:?}: {e}"));
sha256_hex(&bytes)
}
#[allow(dead_code)]
pub fn collect_files_rel(dir: &Path, prefix: &str, out: &mut Vec<String>) {
let entries = std::fs::read_dir(dir).unwrap_or_else(|e| panic!("read_dir {dir:?}: {e}"));
for entry in entries {
let entry = entry.unwrap_or_else(|e| panic!("dir entry under {dir:?}: {e}"));
let name = entry.file_name().to_string_lossy().into_owned();
if name.starts_with("._") || name == ".DS_Store" {
continue;
}
let rel = if prefix.is_empty() {
name
} else {
format!("{prefix}/{name}")
};
let file_type = entry
.file_type()
.unwrap_or_else(|e| panic!("file_type {:?}: {e}", entry.path()));
if file_type.is_dir() {
collect_files_rel(&entry.path(), &rel, out);
} else {
out.push(rel);
}
}
}
#[allow(dead_code)]
pub fn assert_exact_sha_manifest(dir: &Path, cases: &[(&str, &str)]) {
use std::collections::BTreeSet;
let mut found = Vec::new();
collect_files_rel(dir, "", &mut found);
let on_disk: BTreeSet<String> = found.into_iter().collect();
let pinned: BTreeSet<String> = cases.iter().map(|(rel, _)| (*rel).to_owned()).collect();
if on_disk != pinned {
let missing: Vec<&String> = pinned.difference(&on_disk).collect();
let extra: Vec<&String> = on_disk.difference(&pinned).collect();
panic!(
"artifact manifest mismatch under {dir:?}:\n \
missing (pinned but not on disk): {missing:?}\n \
extra (on disk but not pinned): {extra:?}"
);
}
for (relative, expected) in cases {
assert_eq!(
&sha256_file(&dir.join(relative)),
expected,
"sha256 drift on artifact {relative} under {dir:?}"
);
}
}
#[allow(dead_code)]
pub struct Face {
pub id: String,
pub person: String,
pub crop: Vec<u8>,
pub landmarks: [Point; LANDMARK_COUNT],
pub align_matrix: [f64; 6],
pub reference: Vec<f32>,
pub reference_norm: f64,
}
#[allow(dead_code)]
pub struct Reference {
pub faces: Vec<Face>,
pub dim: usize,
pub same_min: f64,
pub different_max: f64,
pub reference_min_same: f64,
pub reference_max_different: f64,
pub worst_same_ids: [String; 2],
pub worst_different_ids: [String; 2],
}
#[allow(dead_code)]
pub fn load_reference() -> Reference {
let faces_dir = fixtures_dir().join("faces");
let manifest = read_json(&faces_dir.join("manifest.json"));
let reference = read_json(&fixtures_dir().join("onnx_reference.json"));
assert_eq!(
reference["source"]["member_sha256"].as_str(),
Some(SOURCE_MEMBER_SHA256),
"the reference was cut from a different ONNX than this kit converts"
);
assert_eq!(
reference["source"]["pack_sha256"].as_str(),
Some(SOURCE_PACK_SHA256),
"the reference's pack pin is not the one the conversion consumed"
);
assert_eq!(
reference["source"]["precision"].as_str(),
Some("fp32"),
"the reference must be fp32; an fp16 oracle would measure the artifact against itself"
);
assert_eq!(reference["preprocessing"]["order"].as_str(), Some("rgb"));
assert_eq!(reference["preprocessing"]["layout"].as_str(), Some("nchw"));
let dim = usize::try_from(reference["dim"].as_u64().expect("`dim`")).expect("dim fits");
let manifest_faces = manifest["faces"].as_array().expect("`faces` array");
let reference_faces = reference["faces"].as_array().expect("`faces` array");
assert_eq!(
manifest_faces.len(),
reference_faces.len(),
"the reference covers a different number of faces than the fixture manifest lists"
);
let faces = manifest_faces
.iter()
.zip(reference_faces)
.map(|(row, reference)| {
let id = row["id"].as_str().expect("fixture id").to_owned();
assert_eq!(
reference["id"].as_str(),
Some(id.as_str()),
"the reference and the fixture manifest disagree on face order"
);
let crop_sha = row["crop_sha256"].as_str().expect("crop sha");
assert_eq!(
reference["crop_sha256"].as_str(),
Some(crop_sha),
"{id}: the reference was cut from a different crop"
);
let crop_path = faces_dir.join(row["crop"].as_str().expect("crop name"));
let crop = std::fs::read(&crop_path).unwrap_or_else(|e| panic!("read {crop_path:?}: {e}"));
assert_eq!(crop.len(), TEMPLATE_BYTES, "{id}: crop length");
assert_eq!(
sha256_hex(&crop),
crop_sha,
"{id}: the committed crop's bytes have moved"
);
let points = row["detection"]["landmarks5"]
.as_array()
.unwrap_or_else(|| panic!("{id}: no landmarks5"));
assert_eq!(points.len(), LANDMARK_COUNT, "{id}: landmark count");
let mut landmarks = [Point::new(0.0, 0.0); LANDMARK_COUNT];
for (slot, point) in landmarks.iter_mut().zip(points) {
let xy = point.as_array().expect("landmark pair");
*slot = Point::new(
xy[0].as_f64().expect("landmark x") as f32,
xy[1].as_f64().expect("landmark y") as f32,
);
}
let rows = row["align_matrix"]
.as_array()
.unwrap_or_else(|| panic!("{id}: no align_matrix"));
let mut align_matrix = [0.0f64; 6];
for (i, matrix_row) in rows.iter().enumerate() {
for (j, value) in matrix_row
.as_array()
.expect("matrix row")
.iter()
.enumerate()
{
align_matrix[i * 3 + j] = value.as_f64().expect("matrix entry");
}
}
let embedding: Vec<f32> = reference["embedding"]
.as_array()
.unwrap_or_else(|| panic!("{id}: no reference embedding"))
.iter()
.map(|v| v.as_f64().expect("finite reference component") as f32)
.collect();
assert_eq!(embedding.len(), dim, "{id}: reference width");
assert!(
embedding.iter().all(|v| v.is_finite()),
"{id}: non-finite reference component"
);
let reference_norm = reference["l2_norm"]
.as_f64()
.unwrap_or_else(|| panic!("{id}: no reference norm"));
Face {
id,
person: row["person"].as_str().expect("person").to_owned(),
crop,
landmarks,
align_matrix,
reference: embedding,
reference_norm,
}
})
.collect();
let pairs = &reference["known_pairs"];
Reference {
faces,
dim,
same_min: pairs["same_min"].as_f64().expect("same_min"),
different_max: pairs["different_max"].as_f64().expect("different_max"),
reference_min_same: pairs["min_same"].as_f64().expect("min_same"),
reference_max_different: pairs["max_different"].as_f64().expect("max_different"),
worst_same_ids: two_ids(&pairs["worst_same_ids"]),
worst_different_ids: two_ids(&pairs["worst_different_ids"]),
}
}
fn two_ids(value: &serde_json::Value) -> [String; 2] {
let ids = value.as_array().expect("id pair");
assert_eq!(ids.len(), 2, "a pair names two faces");
[
ids[0].as_str().expect("id").to_owned(),
ids[1].as_str().expect("id").to_owned(),
]
}
fn read_json(path: &Path) -> serde_json::Value {
let text = std::fs::read_to_string(path).unwrap_or_else(|e| panic!("read {path:?}: {e}"));
serde_json::from_str(&text).unwrap_or_else(|e| panic!("parse {path:?}: {e}"))
}
#[allow(dead_code)]
pub fn solved_transform(face: &Face) -> SimilarityTransform {
SimilarityTransform::estimate(&face.landmarks, &ARCFACE_TEMPLATE)
.unwrap_or_else(|e| panic!("{}: the committed landmarks solve to nothing: {e}", face.id))
}
#[allow(dead_code)]
pub fn cosine(a: &[f32], b: &[f32]) -> f64 {
let dot: f64 = a
.iter()
.zip(b.iter())
.map(|(x, y)| f64::from(*x) * f64::from(*y))
.sum();
let na = a
.iter()
.map(|x| f64::from(*x) * f64::from(*x))
.sum::<f64>()
.sqrt();
let nb = b
.iter()
.map(|x| f64::from(*x) * f64::from(*x))
.sum::<f64>()
.sqrt();
dot / (na * nb)
}