#[path = "../../support/models_lock_manifest.rs"]
#[allow(dead_code)]
pub mod models_lock_manifest;
#[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::audio::ced::CedModel;
#[allow(dead_code)]
pub const VENDOR_DIR: &str = "ced";
#[allow(dead_code)]
pub const ARTIFACT_LOCK_REVISION: &str = "96d504716c4e2ebe6182312f65c4ee6c23921a5d";
#[allow(dead_code)]
pub const TINY_BUNDLE_PATH: &str = "ced-tiny/ced_tiny.mlmodelc";
const TINY_REVISION: &str = "ace276d29dd0bb3f3517b0fa8cf300738c409019";
const MINI_REVISION: &str = "26c3ebcae85d4330f4fc26763f029539a3afcda0";
const SMALL_REVISION: &str = "06bb40c5ec089e96867ebc5246be02441f4a71e4";
const BASE_REVISION: &str = "db3e14a8db4c21b56b165261c39649741a900e7f";
#[allow(dead_code)]
pub const fn ced_revision(model: CedModel) -> &'static str {
match model {
CedModel::Tiny => TINY_REVISION,
CedModel::Mini => MINI_REVISION,
CedModel::Small => SMALL_REVISION,
CedModel::Base => BASE_REVISION,
}
}
#[allow(dead_code)]
pub fn models_dir() -> PathBuf {
std::env::var_os("CED_TEST_MODELS")
.map_or_else(|| workspace_root::models_root().join("ced"), PathBuf::from)
}
#[allow(dead_code)]
pub fn model_root(model: CedModel) -> PathBuf {
models_dir().join(model.dir_name())
}
#[allow(dead_code)]
pub fn model_path(model: CedModel) -> PathBuf {
model.mlmodelc_path(models_dir())
}
#[allow(dead_code)]
pub fn fixture_path(relative: &str) -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("ced")
.join("fixtures")
.join(relative)
}
#[derive(Debug, serde::Deserialize)]
#[allow(dead_code)]
pub struct GoldenClip {
pub id: String,
pub file: String,
pub n_samples: usize,
pub logits: Vec<f32>,
}
#[derive(Debug, serde::Deserialize)]
#[allow(dead_code)]
pub struct OracleProvenance {
pub repo: String,
pub revision: String,
pub file: String,
pub sha256: String,
}
#[derive(Debug, serde::Deserialize)]
#[allow(dead_code)]
pub struct GoldenCorpus {
pub oracle: OracleProvenance,
pub clips: Vec<GoldenClip>,
}
#[allow(dead_code)]
pub fn parse_golden_corpus(bytes: &[u8], model: CedModel) -> GoldenCorpus {
let corpus: GoldenCorpus =
serde_json::from_slice(bytes).unwrap_or_else(|e| panic!("parse golden corpus: {e}"));
assert_eq!(
corpus.oracle.repo,
model.hf_repo(),
"golden corpus oracle repo `{}` does not match {} (`{}`) — cross-size mix-up",
corpus.oracle.repo,
model,
model.hf_repo()
);
corpus
}
#[allow(dead_code)]
pub fn load_golden_corpus(model: CedModel) -> GoldenCorpus {
let path = fixture_path(&format!("goldens/{}/corpus.json", model.as_str()));
let bytes = std::fs::read(&path).unwrap_or_else(|e| panic!("read {}: {e}", path.display()));
parse_golden_corpus(&bytes, model)
}
#[allow(dead_code)]
pub fn read_wav_16k_mono(path: &Path) -> Vec<f32> {
let mut reader =
hound::WavReader::open(path).unwrap_or_else(|e| panic!("open {}: {e}", path.display()));
let spec = reader.spec();
assert_eq!(
spec.sample_rate,
16_000,
"{}: must be 16 kHz",
path.display()
);
assert_eq!(spec.channels, 1, "{}: must be mono", path.display());
match spec.sample_format {
hound::SampleFormat::Int => {
let scale = 1.0 / (1_i64 << (spec.bits_per_sample - 1)) as f32;
reader
.samples::<i32>()
.map(|s| s.unwrap() as f32 * scale)
.collect()
}
hound::SampleFormat::Float => reader.samples::<f32>().map(|s| s.unwrap()).collect(),
}
}
#[allow(dead_code)]
pub fn sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest, Sha256};
Sha256::digest(bytes)
.iter()
.map(|b| format!("{b:02x}"))
.collect()
}
#[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 cosine_checked(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "cosine operands differ in length");
assert!(!a.is_empty(), "cosine of empty vectors is undefined");
let mut dot = 0.0f64;
let mut na = 0.0f64;
let mut nb = 0.0f64;
for (i, (&x, &y)) in a.iter().zip(b.iter()).enumerate() {
assert!(x.is_finite(), "left operand non-finite at {i}");
assert!(y.is_finite(), "right operand non-finite at {i}");
dot += (x as f64) * (y as f64);
na += (x as f64) * (x as f64);
nb += (y as f64) * (y as f64);
}
assert!(na > 0.0 && nb > 0.0, "cosine of a zero vector is undefined");
(dot / (na.sqrt() * nb.sqrt())) as f32
}
#[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: &[(String, String)]) {
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.clone()).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:?}\n \
if the bundle changed intentionally, update the pinned `cases` list AND \
the doc SHA table"
);
}
for (relative, expected) in cases {
let actual = sha256_file(&dir.join(relative));
assert_eq!(
&actual, expected,
"sha256 drift on artifact {relative} under {dir:?}"
);
}
}
#[path = "../../support/model_gate_report.rs"]
mod model_gate_report;
#[test]
fn model_gate_report() {
model_gate_report::report(&[("CED_TEST_MODELS", models_dir())]);
}