use half::f16;
use objc2::AnyThread;
use objc2::rc::Retained;
use objc2::runtime::{AnyObject, ProtocolObject};
use objc2_core_ml::{
MLDictionaryFeatureProvider, MLFeatureProvider, MLFeatureValue, MLModel, MLMultiArray,
MLMultiArrayDataType,
};
use objc2_foundation::{NSArray, NSDictionary, NSNumber, NSString};
use super::ns_error_message;
use crate::runtime::error::RuntimeError;
#[allow(deprecated)]
pub fn predict_f32(
model: &MLModel,
input_name: &str,
mel: &[f32],
shape: &[usize],
output_name: &str,
) -> Result<(Vec<f32>, Vec<usize>), RuntimeError> {
let expected_len: usize = shape.iter().product();
if mel.len() != expected_len {
return Err(RuntimeError::DataLengthMismatch {
expected: expected_len,
got: mel.len(),
});
}
let dims: Vec<Retained<NSNumber>> = shape.iter().map(|&d| NSNumber::new_usize(d)).collect();
let ns_shape: Retained<NSArray<NSNumber>> = NSArray::from_retained_slice(&dims);
let input: Retained<MLMultiArray> = unsafe {
MLMultiArray::initWithShape_dataType_error(
MLMultiArray::alloc(),
&ns_shape,
MLMultiArrayDataType::Float16,
)
}
.map_err(|err| {
RuntimeError::InferenceFailed(format!(
"MLMultiArray init failed: {}",
ns_error_message(&err)
))
})?;
let in_strides = strides_of(&input)?;
{
let base = unsafe { input.dataPointer() }.as_ptr() as *mut f16;
write_strided(base, mel, shape, &in_strides);
}
let feat: Retained<MLFeatureValue> =
unsafe { MLFeatureValue::featureValueWithMultiArray(&input) };
let key = NSString::from_str(input_name);
let value: &AnyObject = &feat;
let dict: Retained<NSDictionary<NSString, AnyObject>> =
NSDictionary::from_slices(&[&*key], &[value]);
let provider: Retained<MLDictionaryFeatureProvider> = unsafe {
MLDictionaryFeatureProvider::initWithDictionary_error(
MLDictionaryFeatureProvider::alloc(),
&dict,
)
}
.map_err(|err| {
RuntimeError::InferenceFailed(format!(
"feature provider init failed: {}",
ns_error_message(&err)
))
})?;
let provider_obj: &ProtocolObject<dyn MLFeatureProvider> = ProtocolObject::from_ref(&*provider);
let result: Retained<ProtocolObject<dyn MLFeatureProvider>> =
unsafe { model.predictionFromFeatures_error(provider_obj) }.map_err(|err| {
RuntimeError::InferenceFailed(format!("prediction failed: {}", ns_error_message(&err)))
})?;
let out_key = NSString::from_str(output_name);
let out_feat: Retained<MLFeatureValue> = unsafe { result.featureValueForName(&out_key) }
.ok_or_else(|| {
RuntimeError::InferenceFailed(format!("output '{output_name}' missing from result"))
})?;
let out_arr: Retained<MLMultiArray> =
unsafe { out_feat.multiArrayValue() }.ok_or_else(|| {
RuntimeError::InferenceFailed(format!("output '{output_name}' is not a multi-array"))
})?;
let out_shape = shape_of(&out_arr)?;
let out_strides = strides_of(&out_arr)?;
let out_len: usize = out_shape.iter().product();
let out_dtype = unsafe { out_arr.dataType() };
let raw = unsafe { out_arr.dataPointer() }.as_ptr();
let data = match out_dtype {
MLMultiArrayDataType::Float16 => {
read_strided_f16(raw as *const f16, &out_shape, &out_strides)
}
MLMultiArrayDataType::Float32 => {
read_strided_f32(raw as *const f32, &out_shape, &out_strides)
}
other => {
return Err(RuntimeError::InferenceFailed(format!(
"unsupported output dataType {other:?}"
)));
}
};
debug_assert_eq!(data.len(), out_len);
Ok((data, out_shape))
}
fn shape_of(arr: &MLMultiArray) -> Result<Vec<usize>, RuntimeError> {
let ns: Retained<NSArray<NSNumber>> = unsafe { arr.shape() };
Ok(nsarray_usize(&ns))
}
fn strides_of(arr: &MLMultiArray) -> Result<Vec<usize>, RuntimeError> {
let ns: Retained<NSArray<NSNumber>> = unsafe { arr.strides() };
Ok(nsarray_usize(&ns))
}
fn nsarray_usize(ns: &NSArray<NSNumber>) -> Vec<usize> {
let n = ns.count();
let mut out = Vec::with_capacity(n);
for i in 0..n {
let num = ns.objectAtIndex(i);
out.push(num.as_usize());
}
out
}
fn write_strided(base: *mut f16, data: &[f32], shape: &[usize], strides: &[usize]) {
let rank = shape.len();
let total = data.len();
let mut idx = vec![0usize; rank];
for &v in data.iter().take(total) {
let mut off = 0usize;
for d in 0..rank {
off += idx[d] * strides[d];
}
unsafe { *base.add(off) = f16::from_f32(v) };
for d in (0..rank).rev() {
idx[d] += 1;
if idx[d] < shape[d] {
break;
}
idx[d] = 0;
}
}
}
fn read_strided_f16(base: *const f16, shape: &[usize], strides: &[usize]) -> Vec<f32> {
read_strided_with(shape, strides, |off| unsafe { (*base.add(off)).to_f32() })
}
fn read_strided_f32(base: *const f32, shape: &[usize], strides: &[usize]) -> Vec<f32> {
read_strided_with(shape, strides, |off| unsafe { *base.add(off) })
}
fn read_strided_with(
shape: &[usize],
strides: &[usize],
mut read: impl FnMut(usize) -> f32,
) -> Vec<f32> {
let rank = shape.len();
let total: usize = shape.iter().product();
let mut out = Vec::with_capacity(total);
let mut idx = vec![0usize; rank];
for _ in 0..total {
let mut off = 0usize;
for d in 0..rank {
off += idx[d] * strides[d];
}
out.push(read(off));
for d in (0..rank).rev() {
idx[d] += 1;
if idx[d] < shape[d] {
break;
}
idx[d] = 0;
}
}
out
}