#![allow(missing_docs)]
#[path = "../../tests/support/workspace_root.rs"]
#[allow(dead_code)]
mod workspace_root;
use core::{sync::atomic::AtomicBool, time::Duration};
use std::{
hint::black_box,
path::{Path, PathBuf},
};
use coremlit::{
Model, MultiArray,
audio::align::{
ANALYSIS_TIMEBASE, Aligner, EnglishNormalizer, Lang, OutputClock, default_oov_policy,
encode::{DEFAULT_ENCODER_COMPUTE, ENCODER_WINDOW_SAMPLES},
},
};
use criterion::{Criterion, criterion_group, criterion_main};
const JFK_TRANSCRIPT: &str = "And so my fellow Americans ask not what your country can do for you, \
ask what you can do for your country.";
const JFK_SECONDS: f64 = 11.0;
fn models_dir() -> PathBuf {
std::env::var_os("ALIGNKIT_TEST_MODELS").map_or_else(
|| workspace_root::models_root().join("alignkit"),
PathBuf::from,
)
}
fn load_wav_mono_f32(path: &Path) -> Vec<f32> {
let mut reader = hound::WavReader::open(path).expect("fixture wav opens");
let spec = reader.spec();
assert_eq!(spec.channels, 1, "fixture must be mono");
assert_eq!(spec.sample_rate, 16_000, "fixture must be 16 kHz");
assert_eq!(spec.sample_format, hound::SampleFormat::Int);
reader
.samples::<i16>()
.map(|s| f32::from(s.expect("valid sample")) / 32_768.0)
.collect()
}
fn bench_align(c: &mut Criterion) {
let model = models_dir().join("base960h_aligner.mlmodelc");
let samples = load_wav_mono_f32(
&PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/whisper/fixtures/audio/jfk.wav"),
);
let aligner = Aligner::from_paths(Lang::En, &model, Box::new(EnglishNormalizer::new()))
.expect("build the En aligner (set ALIGNKIT_TEST_MODELS to the model directory)");
let encoder = Model::load(&model, DEFAULT_ENCODER_COMPUTE).expect("load the CoreML encoder");
let mut window = vec![0.0f32; ENCODER_WINDOW_SAMPLES];
window[..samples.len()].copy_from_slice(&samples);
let waveform =
MultiArray::from_slice(&[1, ENCODER_WINDOW_SAMPLES], &window).expect("the input window");
let clock = OutputClock::new(0, ANALYSIS_TIMEBASE, 0).expect("clock construction");
let abort = AtomicBool::new(false);
let mut group = c.benchmark_group("alignkit");
group
.sample_size(10)
.measurement_time(Duration::from_secs(30))
.throughput(criterion::Throughput::Elements(JFK_SECONDS as u64));
group.bench_function("encode", |b| {
b.iter(|| {
black_box(
encoder
.predict_with(&[("waveform", black_box(&waveform))])
.expect("encode"),
);
});
});
group.bench_function("align_chunk", |b| {
b.iter_batched(
|| {
aligner
.detect_oov(JFK_TRANSCRIPT)
.expect("detect_oov")
.decide(default_oov_policy)
},
|resolution| {
black_box(
aligner
.align_chunk(
black_box(&samples),
&[],
black_box(JFK_TRANSCRIPT),
clock,
&abort,
resolution,
)
.expect("align_chunk"),
);
},
criterion::BatchSize::SmallInput,
);
});
group.finish();
println!(
"\nalignment RTF = (wall time reported above) / {JFK_SECONDS:.1} s of audio; speed factor = \
its reciprocal. ASR runtime is NOT included — no ASR model is loaded by this bench.\n"
);
}
criterion_group!(benches, bench_align);
criterion_main!(benches);