use objc2::{
AnyThread,
rc::Retained,
runtime::{AnyObject, ProtocolObject},
};
use objc2_core_ml::{MLDictionaryFeatureProvider, MLFeatureProvider, MLFeatureValue};
use objc2_foundation::{NSDictionary, NSString};
use crate::{MultiArray, NsErrorInfo, PredictionError};
#[derive(Debug)]
pub struct Features {
entries: Vec<(String, MultiArray)>,
}
impl Default for Features {
fn default() -> Self {
Self::new()
}
}
impl Features {
#[inline(always)]
pub const fn new() -> Self {
Self {
entries: Vec::new(),
}
}
pub fn insert(&mut self, name: impl Into<String>, array: MultiArray) -> &mut Self {
let name = name.into();
self.entries.retain(|(existing, _)| *existing != name);
self.entries.push((name, array));
self
}
#[must_use]
pub fn with(mut self, name: impl Into<String>, array: MultiArray) -> Self {
self.insert(name, array);
self
}
pub fn get(&self, name: &str) -> Option<&MultiArray> {
self.entries.iter().find(|(n, _)| n == name).map(|(_, a)| a)
}
pub fn take(&mut self, name: &str) -> Option<MultiArray> {
let index = self.entries.iter().position(|(n, _)| n == name)?;
Some(self.entries.remove(index).1)
}
pub fn names(&self) -> impl Iterator<Item = &str> {
self.entries.iter().map(|(n, _)| n.as_str())
}
#[inline(always)]
pub const fn len(&self) -> usize {
self.entries.len()
}
#[inline(always)]
pub const fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub(crate) fn byte_ranges(&self) -> Vec<(usize, usize)> {
self.entries.iter().map(|(_, a)| a.byte_range()).collect()
}
pub(crate) fn to_provider(
&self,
) -> Result<Retained<MLDictionaryFeatureProvider>, PredictionError> {
provider_from_pairs(
self
.entries
.iter()
.map(|(name, array)| (name.as_str(), array)),
)
}
pub(crate) fn from_provider(
provider: &ProtocolObject<dyn MLFeatureProvider>,
wanted: Option<&[&str]>,
known_regions: &mut Vec<(usize, usize)>,
) -> Result<Self, PredictionError> {
fn overlaps(a: (usize, usize), b: (usize, usize)) -> bool {
a.0 < b.1 && b.0 < a.1
}
let mut features = Self::new();
let names = unsafe { provider.featureNames() };
for name in names.iter() {
let name_str = name.to_string();
if wanted.is_some_and(|wanted| !wanted.contains(&name_str.as_str())) {
continue;
}
let value = unsafe { provider.featureValueForName(&name) }
.ok_or_else(|| PredictionError::MissingOutput(name_str.clone()))?;
let array = unsafe { value.multiArrayValue() }
.ok_or_else(|| PredictionError::NotMultiArray(name_str.clone()))?;
let mut array = MultiArray::from_raw(array);
let region = array.byte_range();
if known_regions.iter().any(|&known| overlaps(known, region)) {
array = array
.deep_copy()
.map_err(PredictionError::AliasCopyFailed)?;
}
known_regions.push(array.byte_range());
features.insert(name_str, array);
}
Ok(features)
}
}
pub(crate) fn provider_from_pairs<'a, I>(
pairs: I,
) -> Result<Retained<MLDictionaryFeatureProvider>, PredictionError>
where
I: Iterator<Item = (&'a str, &'a MultiArray)>,
{
let (lower, _) = pairs.size_hint();
let mut keys: Vec<Retained<NSString>> = Vec::with_capacity(lower);
let mut values: Vec<Retained<AnyObject>> = Vec::with_capacity(lower);
for (name, array) in pairs {
keys.push(NSString::from_str(name));
let value: Retained<MLFeatureValue> =
unsafe { MLFeatureValue::featureValueWithMultiArray(array.raw()) };
values.push(value.into());
}
let key_refs: Vec<&NSString> = keys.iter().map(|k| k.as_ref()).collect();
let dict = NSDictionary::from_retained_objects(&key_refs, &values);
unsafe {
MLDictionaryFeatureProvider::initWithDictionary_error(
MLDictionaryFeatureProvider::alloc(),
&dict,
)
}
.map_err(|e| PredictionError::Native(NsErrorInfo::from_ns_error(&e)))
}
#[cfg(test)]
mod tests;