mod common;
use coremlit::{
ComputeUnits,
embeddings::clap::{
AudioEncoder, AudioEncoderOptions, Embedding, TextEncoder, TextEncoderOptions,
},
};
const AUDIO_MIN_COSINE: f32 = 0.9999;
const TEXT_MIN_COSINE: f32 = 0.9999;
const AUDIO_INT8_MIN_COSINE: f32 = 0.9999;
const TEXT_INT8_MIN_COSINE: f32 = 0.9999;
const AUDIO_UNITS: &[ComputeUnits] = &[
ComputeUnits::All,
ComputeUnits::CpuAndNeuralEngine,
ComputeUnits::CpuAndGpu,
ComputeUnits::CpuOnly,
];
const TEXT_UNITS: &[ComputeUnits] = &[
ComputeUnits::All,
ComputeUnits::CpuAndNeuralEngine,
ComputeUnits::CpuAndGpu,
ComputeUnits::CpuOnly,
];
fn assert_band(label: &str, unit: ComputeUnits, cos: f32, min: f32) {
assert!(
(min..=1.0 + 1e-6).contains(&cos),
"{label} [{}] cosine vs CpuOnly = {cos:.8} outside [{min}, 1.0] — a placement changed the \
numerics materially (a finding, not a threshold to loosen)",
unit.as_str()
);
}
#[test]
#[ignore = "requires local clapkit models (CLAPKIT_TEST_MODELS)"]
fn audio_placement_agreement_characterized() {
let samples = common::deterministic_window(coremlit::embeddings::clap::audio::TARGET_SAMPLES);
let embed = |unit: ComputeUnits| -> Embedding {
AudioEncoder::from_file_with(
common::audio_model_path(),
AudioEncoderOptions::new().with_compute(unit),
)
.unwrap_or_else(|e| panic!("load audio [{}]: {e}", unit.as_str()))
.embed_window(&samples)
.unwrap_or_else(|e| panic!("embed audio [{}]: {e}", unit.as_str()))
};
let reference = embed(ComputeUnits::CpuOnly);
let mut worst = 1.0f32;
for &unit in AUDIO_UNITS {
let cos = embed(unit).cosine(&reference);
worst = worst.min(cos);
assert_band("audio", unit, cos, AUDIO_MIN_COSINE);
}
assert!((reference.cosine(&reference) - 1.0).abs() <= 1e-5);
eprintln!("[placement] audio worst cross-unit cosine = {worst:.8}");
}
#[test]
#[ignore = "requires local clapkit models (CLAPKIT_TEST_MODELS)"]
fn text_placement_agreement_characterized() {
const PROMPT: &str = "a violin playing a slow melody in a concert hall";
let embed = |unit: ComputeUnits| -> Embedding {
TextEncoder::from_bundled_tokenizer(
common::text_model_path(),
TextEncoderOptions::new().with_compute(unit),
)
.unwrap_or_else(|e| panic!("load text [{}]: {e}", unit.as_str()))
.embed(PROMPT)
.unwrap_or_else(|e| panic!("embed text [{}]: {e}", unit.as_str()))
};
let reference = embed(ComputeUnits::CpuOnly);
let mut worst = 1.0f32;
for &unit in TEXT_UNITS {
let cos = embed(unit).cosine(&reference);
worst = worst.min(cos);
assert_band("text", unit, cos, TEXT_MIN_COSINE);
}
assert!((reference.cosine(&reference) - 1.0).abs() <= 1e-5);
eprintln!("[placement] text worst cross-unit cosine = {worst:.8}");
}
#[test]
#[ignore = "requires local clapkit int8 models (CLAPKIT_TEST_MODELS)"]
fn audio_int8_placement_agreement_characterized() {
let samples = common::deterministic_window(coremlit::embeddings::clap::audio::TARGET_SAMPLES);
let embed = |unit: ComputeUnits| -> Embedding {
AudioEncoder::from_file_with(
common::audio_model_int8_path(),
AudioEncoderOptions::new().with_compute(unit),
)
.unwrap_or_else(|e| panic!("load audio int8 [{}]: {e}", unit.as_str()))
.embed_window(&samples)
.unwrap_or_else(|e| panic!("embed audio int8 [{}]: {e}", unit.as_str()))
};
let reference = embed(ComputeUnits::CpuOnly);
let mut worst = 1.0f32;
for &unit in AUDIO_UNITS {
let cos = embed(unit).cosine(&reference);
worst = worst.min(cos);
assert_band("audio int8", unit, cos, AUDIO_INT8_MIN_COSINE);
}
assert!((reference.cosine(&reference) - 1.0).abs() <= 1e-5);
eprintln!("[placement] audio int8 worst cross-unit cosine = {worst:.8}");
}
#[test]
#[ignore = "requires local clapkit int8 models (CLAPKIT_TEST_MODELS)"]
fn text_int8_placement_agreement_characterized() {
const PROMPT: &str = "a violin playing a slow melody in a concert hall";
let embed = |unit: ComputeUnits| -> Embedding {
TextEncoder::from_bundled_tokenizer(
common::text_model_int8_path(),
TextEncoderOptions::new().with_compute(unit),
)
.unwrap_or_else(|e| panic!("load text int8 [{}]: {e}", unit.as_str()))
.embed(PROMPT)
.unwrap_or_else(|e| panic!("embed text int8 [{}]: {e}", unit.as_str()))
};
let reference = embed(ComputeUnits::CpuOnly);
let mut worst = 1.0f32;
for &unit in TEXT_UNITS {
let cos = embed(unit).cosine(&reference);
worst = worst.min(cos);
assert_band("text int8", unit, cos, TEXT_INT8_MIN_COSINE);
}
assert!((reference.cosine(&reference) - 1.0).abs() <= 1e-5);
eprintln!("[placement] text int8 worst cross-unit cosine = {worst:.8}");
}