use std::sync::Arc;
use objc2::rc::Retained;
use objc2_core_ml::MLModel;
use crate::runtime::{
error::RuntimeError,
session::RuntimeSession,
tensor::{Shape, Tensor, TensorData, TensorView},
};
use super::bridge;
const ENC_DIM: usize = 768;
const N_MELS: usize = 64;
const FILL_FLOOR: f64 = 0.5;
pub struct SharedModel(pub Retained<MLModel>);
unsafe impl Send for SharedModel {}
unsafe impl Sync for SharedModel {}
pub struct BucketModel {
pub size: usize,
pub model: Arc<SharedModel>,
}
pub fn calc_output_length(t: usize) -> usize {
fn stride2(x: usize) -> usize {
(x - 1) / 2 + 1
}
stride2(stride2(t))
}
pub fn select_bucket(t: usize, buckets: &[usize], floor: f64) -> Option<usize> {
buckets
.iter()
.copied()
.filter(|&n| n >= t && t as f64 / n as f64 >= floor)
.min()
}
pub fn pad_time(mel: &[f32], channels: usize, t: usize, n: usize) -> Vec<f32> {
debug_assert!(t <= n, "pad target {n} must be >= source {t}");
debug_assert_eq!(mel.len(), channels * t, "mel len mismatch in pad_time");
let mut out = vec![0.0f32; channels * n];
for c in 0..channels {
let src = &mel[c * t..c * t + t];
out[c * n..c * n + t].copy_from_slice(src);
}
out
}
pub fn trim_time(out: &[f32], channels: usize, t_padded: usize, t_keep: usize) -> Vec<f32> {
debug_assert!(t_keep <= t_padded, "trim {t_keep} must be <= {t_padded}");
debug_assert_eq!(
out.len(),
channels * t_padded,
"out len mismatch in trim_time"
);
let mut trimmed = Vec::with_capacity(channels * t_keep);
for c in 0..channels {
let row = &out[c * t_padded..c * t_padded + t_keep];
trimmed.extend_from_slice(row);
}
trimmed
}
pub struct AneEncoderSession {
buckets: Vec<BucketModel>,
ort_fallback: Box<dyn RuntimeSession>,
}
impl AneEncoderSession {
pub fn new(mut buckets: Vec<BucketModel>, ort_fallback: Box<dyn RuntimeSession>) -> Self {
buckets.sort_by_key(|b| b.size);
Self {
buckets,
ort_fallback,
}
}
fn bucket_sizes(&self) -> Vec<usize> {
self.buckets.iter().map(|b| b.size).collect()
}
fn model_for(&self, size: usize) -> Option<&Arc<SharedModel>> {
self.buckets
.iter()
.find(|b| b.size == size)
.map(|b| &b.model)
}
}
impl RuntimeSession for AneEncoderSession {
fn run(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
if inputs.is_empty() {
return Err(RuntimeError::InvalidInputCount {
expected: 2,
got: inputs.len(),
});
}
let mel_view: TensorView<'_> = inputs[0].view();
let mel = mel_view.data().as_f32().ok_or_else(|| {
RuntimeError::InferenceFailed("ANE encoder mel input is not f32".to_string())
})?;
let dims = mel_view.shape().dims();
if dims.len() != 3 || dims[0] != 1 || dims[1] != N_MELS {
return Err(RuntimeError::InferenceFailed(format!(
"ANE encoder expects mel shape [1, {N_MELS}, T], got {dims:?}"
)));
}
let t = dims[2];
let sizes = self.bucket_sizes();
match select_bucket(t, &sizes, FILL_FLOOR) {
Some(n) => {
let model = self.model_for(n).ok_or_else(|| {
RuntimeError::InferenceFailed(format!("no compiled model for bucket {n}"))
})?;
let padded = pad_time(mel, N_MELS, t, n);
let (out, out_shape) =
bridge::predict_f32(&model.0, "mel", &padded, &[1, N_MELS, n], "encoded")?;
if out_shape.len() != 3 || out_shape[0] != 1 || out_shape[1] != ENC_DIM {
return Err(RuntimeError::InferenceFailed(format!(
"ANE encoder output shape {out_shape:?} != [1, {ENC_DIM}, T']"
)));
}
let t_padded_prime = out_shape[2];
let t_real_prime = calc_output_length(t);
if t_real_prime > t_padded_prime {
return Err(RuntimeError::InferenceFailed(format!(
"computed encoder output length {t_real_prime} exceeds bucket output {t_padded_prime}"
)));
}
tracing::debug!(
t,
bucket = n,
t_real_prime,
t_padded_prime,
"ANE encoder path (bucketed pad-up)"
);
let trimmed = trim_time(&out, ENC_DIM, t_padded_prime, t_real_prime);
Ok(vec![
Tensor::new(
Shape::new(vec![1, ENC_DIM, t_real_prime]),
TensorData::F32(trimmed),
)?,
Tensor::new(
Shape::new(vec![1]),
TensorData::I64(vec![t_real_prime as i64]),
)?,
])
}
None => {
tracing::debug!(
t,
"ANE encoder path (ort fallback: no bucket within fill-floor)"
);
self.ort_fallback.run(inputs)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_calc_output_length_known_points() {
assert_eq!(calc_output_length(768), 192);
assert_eq!(calc_output_length(1500), 375);
assert_eq!(calc_output_length(3000), 750);
assert_eq!(calc_output_length(593), 149);
assert_eq!(calc_output_length(250), 63);
assert_eq!(calc_output_length(769), 193);
assert_eq!(calc_output_length(400), 100);
}
#[test]
fn test_select_bucket_fill_floor_cases() {
let buckets = [768usize, 1536, 3000];
assert_eq!(select_bucket(400, &buckets, 0.5), Some(768));
assert_eq!(select_bucket(200, &buckets, 0.5), None);
assert_eq!(select_bucket(800, &buckets, 0.5), Some(1536));
assert_eq!(select_bucket(769, &buckets, 0.5), Some(1536));
assert_eq!(select_bucket(4000, &buckets, 0.5), None);
assert_eq!(select_bucket(768, &buckets, 0.5), Some(768));
}
#[test]
fn test_select_bucket_fill_equal_to_floor_is_selected() {
assert_eq!(select_bucket(384, &[768], 0.5), Some(768));
}
#[test]
fn test_select_bucket_unsorted_input() {
let buckets = [3000usize, 768, 1536];
assert_eq!(select_bucket(400, &buckets, 0.5), Some(768));
assert_eq!(select_bucket(800, &buckets, 0.5), Some(1536));
}
#[test]
fn test_pad_time_appends_zeros_per_channel() {
let mel = vec![1.0, 2.0, 3.0, 4.0];
let padded = pad_time(&mel, 2, 2, 4);
assert_eq!(padded, vec![1.0, 2.0, 0.0, 0.0, 3.0, 4.0, 0.0, 0.0]);
assert_eq!(padded.len(), 2 * 4);
}
#[test]
fn test_pad_time_noop_when_t_equals_n() {
let mel = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let padded = pad_time(&mel, 2, 3, 3);
assert_eq!(padded, mel);
}
#[test]
fn test_trim_time_keeps_leading_frames_per_channel() {
let out = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let trimmed = trim_time(&out, 2, 4, 2);
assert_eq!(trimmed, vec![1.0, 2.0, 5.0, 6.0]);
assert_eq!(trimmed.len(), 2 * 2);
}
#[test]
fn test_pad_then_trim_roundtrip_recovers_prefix() {
let mel = vec![10.0, 20.0, 30.0, 40.0, 50.0, 60.0]; let padded = pad_time(&mel, 2, 3, 5);
let back = trim_time(&padded, 2, 5, 3);
assert_eq!(back, mel);
}
#[test]
fn test_streaming_window_routes_to_ort_fallback() {
use crate::runtime::mock::MockSession;
const SHIPPED_BUCKETS: &[usize] = &[768, 1536, 3000];
const T: usize = 250;
assert_eq!(
select_bucket(T, SHIPPED_BUCKETS, FILL_FLOOR),
None,
"a 250-frame streaming window must not select any shipped bucket"
);
let t_prime = calc_output_length(T);
let fallback = MockSession::new(
vec![Shape::new(vec![1, N_MELS, T]), Shape::new(vec![1])],
vec![
Tensor::new(
Shape::new(vec![1, ENC_DIM, t_prime]),
TensorData::F32(vec![0.0; ENC_DIM * t_prime]),
)
.unwrap(),
Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![t_prime as i64])).unwrap(),
],
);
let session = AneEncoderSession::new(Vec::new(), Box::new(fallback));
let mel = Tensor::new(
Shape::new(vec![1, N_MELS, T]),
TensorData::F32(vec![0.0; N_MELS * T]),
)
.unwrap();
let len = Tensor::new(Shape::new(vec![1]), TensorData::I64(vec![T as i64])).unwrap();
let out = session.run(&[mel, len]).expect("fallback run succeeds");
assert_eq!(out.len(), 2, "encoder emits [encoded, encoded_len]");
assert_eq!(out[0].shape().dims(), &[1, ENC_DIM, t_prime]);
match out[1].view().data() {
crate::runtime::tensor::TensorDataView::I64(v) => {
assert_eq!(v[0], t_prime as i64, "fallback encoded_len passes through")
}
other => panic!("expected I64 encoded_len, got {other:?}"),
}
}
}