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;
const STREAMING_FILL_FLOOR: f64 = 0.0;
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)
}
fn run_with_floor(&self, inputs: &[Tensor], floor: f64) -> 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, 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,
floor,
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,
floor,
"ANE encoder path (ort fallback: no bucket within fill-floor)"
);
self.ort_fallback.run(inputs)
}
}
}
}
impl RuntimeSession for AneEncoderSession {
fn run(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
self.run_with_floor(inputs, FILL_FLOOR)
}
fn run_low_latency(&self, inputs: &[Tensor]) -> Result<Vec<Tensor>, RuntimeError> {
self.run_with_floor(inputs, STREAMING_FILL_FLOOR)
}
fn is_ane_encoder(&self) -> bool {
true
}
}
#[cfg(test)]
mod tests;