use super::*;
use crate::{DataType, MultiArray};
#[test]
fn insert_get_take_names() {
let mut features = Features::new();
features
.insert("audio", MultiArray::zeros(&[4], DataType::F32).unwrap())
.insert("mask", MultiArray::zeros(&[2], DataType::F32).unwrap());
assert_eq!(features.len(), 2);
assert_eq!(features.names().collect::<Vec<_>>(), vec!["audio", "mask"]);
assert_eq!(features.get("audio").unwrap().count(), 4);
assert_eq!(features.take("mask").unwrap().count(), 2);
assert!(features.get("mask").is_none());
}
#[test]
fn provider_round_trip_preserves_names_shapes_and_data() {
let mut features = Features::new();
features.insert(
"x",
MultiArray::from_slice(&[2, 2], &[1.0f32, 2.0, 3.0, 4.0]).unwrap(),
);
let provider = features.to_provider().unwrap();
let mut known_regions = Vec::new();
let back = Features::from_provider(
ProtocolObject::from_ref(&*provider),
None,
&mut known_regions,
)
.unwrap();
let x = back.get("x").unwrap();
assert_eq!(x.shape(), vec![2, 2]);
assert_eq!(x.as_slice::<f32>().unwrap(), &[1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn from_provider_deep_copies_one_array_shared_under_two_names() {
let shared = MultiArray::from_slice(&[2, 2], &[1.0f32, 2.0, 3.0, 4.0]).unwrap();
let value: Retained<MLFeatureValue> =
unsafe { MLFeatureValue::featureValueWithMultiArray(shared.raw()) };
let value: Retained<AnyObject> = value.into();
let keys = [NSString::from_str("a"), NSString::from_str("b")];
let key_refs: Vec<&NSString> = keys.iter().map(AsRef::as_ref).collect();
let dict = NSDictionary::from_retained_objects(&key_refs, &[value.clone(), value]);
let provider = unsafe {
MLDictionaryFeatureProvider::initWithDictionary_error(
MLDictionaryFeatureProvider::alloc(),
&dict,
)
}
.unwrap();
let mut known_regions = Vec::new();
let mut extracted = Features::from_provider(
ProtocolObject::from_ref(&*provider),
None,
&mut known_regions,
)
.unwrap();
let a = extracted.get("a").unwrap();
let b = extracted.get("b").unwrap();
assert_ne!(a.byte_range().0, b.byte_range().0);
assert_eq!(a.as_slice::<f32>().unwrap(), b.as_slice::<f32>().unwrap());
let a_owned = extracted.take("a").unwrap();
let mut b_owned = extracted.take("b").unwrap();
b_owned.as_slice_mut::<f32>().unwrap()[0] = 99.0;
assert_eq!(a_owned.as_slice::<f32>().unwrap()[0], 1.0);
}
#[test]
fn from_provider_deep_copies_output_that_aliases_a_seeded_input() {
let input = MultiArray::from_slice(&[2], &[5.0f32, 6.0]).unwrap();
let input_region = input.byte_range();
let value: Retained<MLFeatureValue> =
unsafe { MLFeatureValue::featureValueWithMultiArray(input.raw()) };
let value: Retained<AnyObject> = value.into();
let key = NSString::from_str("y");
let key_refs: Vec<&NSString> = vec![key.as_ref()];
let dict = NSDictionary::from_retained_objects(&key_refs, &[value]);
let provider = unsafe {
MLDictionaryFeatureProvider::initWithDictionary_error(
MLDictionaryFeatureProvider::alloc(),
&dict,
)
}
.unwrap();
let mut known_regions = vec![input_region];
let extracted = Features::from_provider(
ProtocolObject::from_ref(&*provider),
None,
&mut known_regions,
)
.unwrap();
let output = extracted.get("y").unwrap();
assert_ne!(output.byte_range().0, input_region.0);
assert_eq!(output.as_slice::<f32>().unwrap(), &[5.0, 6.0]);
}
#[test]
fn insert_replacing_moves_name_to_end() {
let mut features = Features::new();
features
.insert("a", MultiArray::zeros(&[1], DataType::F32).unwrap())
.insert("b", MultiArray::zeros(&[2], DataType::F32).unwrap());
features.insert("a", MultiArray::zeros(&[9], DataType::F32).unwrap());
assert_eq!(features.names().collect::<Vec<_>>(), vec!["b", "a"]);
assert_eq!(features.get("a").unwrap().count(), 9);
}
#[test]
fn overlapping_offset_regions_are_detected() {
let base = MultiArray::from_slice(&[8], &[0.0f32; 8]).unwrap();
let (start, end) = base.byte_range();
assert_eq!(end - start, 8 * 4);
let mut known = vec![(start, end)];
let offset_view = (start + 8, start + 16); assert!(
known
.iter()
.any(|&k| k.0 < offset_view.1 && offset_view.0 < k.1)
);
let adjacent = (end, end + 16); assert!(!known.iter().any(|&k| k.0 < adjacent.1 && adjacent.0 < k.1));
known.push(adjacent);
}
#[test]
fn provider_from_borrowed_pairs_round_trips() {
let x = MultiArray::from_slice(&[2], &[1.0f32, 2.0]).unwrap();
let y = MultiArray::from_slice(&[1], &[3.0f32]).unwrap();
let provider = super::provider_from_pairs([("x", &x), ("y", &y)].into_iter()).unwrap();
let mut known_regions = Vec::new();
let back = Features::from_provider(
ProtocolObject::from_ref(&*provider),
None,
&mut known_regions,
)
.unwrap();
assert_eq!(
back.get("x").unwrap().as_slice::<f32>().unwrap(),
&[1.0, 2.0]
);
assert_eq!(back.get("y").unwrap().as_slice::<f32>().unwrap(), &[3.0]);
assert_eq!(x.count() + y.count(), 3);
}
#[test]
fn from_provider_materialises_only_the_named_features() {
use objc2::{AllocAnyThread, runtime::ProtocolObject};
use objc2_core_ml::{MLDictionaryFeatureProvider, MLFeatureValue};
use objc2_foundation::{NSDictionary, NSString};
let embedding = MultiArray::from_slice(&[1, 2], &[1.0f32, 2.0]).unwrap();
let provider = unsafe {
let tensor: Retained<AnyObject> =
MLFeatureValue::featureValueWithMultiArray(embedding.raw()).into();
let label: Retained<AnyObject> =
MLFeatureValue::featureValueWithString(&NSString::from_str("speech")).into();
let keys = [NSString::from_str("embedding"), NSString::from_str("label")];
let key_refs: Vec<&NSString> = keys.iter().map(AsRef::as_ref).collect();
let dict = NSDictionary::from_retained_objects(&key_refs, &[tensor, label]);
MLDictionaryFeatureProvider::initWithDictionary_error(
MLDictionaryFeatureProvider::alloc(),
&dict,
)
.expect("a tensor beside a string is a valid feature dictionary")
};
let extracted = Features::from_provider(
ProtocolObject::from_ref(&*provider),
Some(&["embedding"]),
&mut Vec::new(),
)
.expect("an output nobody named must not decide whether this succeeds");
assert_eq!(extracted.names().collect::<Vec<_>>(), vec!["embedding"]);
assert_eq!(
extracted
.get("embedding")
.unwrap()
.as_slice::<f32>()
.unwrap(),
&[1.0, 2.0]
);
let error = Features::from_provider(ProtocolObject::from_ref(&*provider), None, &mut Vec::new())
.unwrap_err();
assert_eq!(error, crate::PredictionError::NotMultiArray("label".into()));
let extracted = Features::from_provider(
ProtocolObject::from_ref(&*provider),
Some(&["embedding", "absent"]),
&mut Vec::new(),
)
.expect("a name the model does not produce is not a failure here");
assert_eq!(extracted.names().collect::<Vec<_>>(), vec!["embedding"]);
}
#[test]
fn from_provider_rejects_non_multi_array_values() {
use objc2::{AllocAnyThread, runtime::ProtocolObject};
use objc2_core_ml::{MLDictionaryFeatureProvider, MLFeatureValue};
use objc2_foundation::{NSDictionary, NSString};
let provider = unsafe {
let value: Retained<AnyObject> =
MLFeatureValue::featureValueWithString(&NSString::from_str("not a tensor")).into();
let dict = NSDictionary::from_retained_objects(&[&*NSString::from_str("meta")], &[value]);
MLDictionaryFeatureProvider::initWithDictionary_error(
MLDictionaryFeatureProvider::alloc(),
&dict,
)
.expect("string feature dictionary is valid")
};
let err = Features::from_provider(ProtocolObject::from_ref(&*provider), None, &mut Vec::new())
.unwrap_err();
assert_eq!(err, crate::PredictionError::NotMultiArray("meta".into()));
}
use objc2_foundation::NSObjectProtocol;
objc2::define_class!(
#[unsafe(super(objc2_foundation::NSObject))]
#[name = "CoremlitGhostFeatureProvider"]
struct GhostProvider;
unsafe impl NSObjectProtocol for GhostProvider {}
unsafe impl MLFeatureProvider for GhostProvider {
#[unsafe(method_id(featureNames))]
fn feature_names(&self) -> Retained<objc2_foundation::NSSet<NSString>> {
objc2_foundation::NSSet::from_retained_slice(&[NSString::from_str("ghost")])
}
#[unsafe(method_id(featureValueForName:))]
fn feature_value_for_name(&self, _name: &NSString) -> Option<Retained<MLFeatureValue>> {
None
}
}
);
#[test]
fn from_provider_surfaces_missing_outputs() {
use objc2::{AllocAnyThread, rc::Retained, runtime::ProtocolObject};
let provider: Retained<GhostProvider> = unsafe { objc2::msg_send![GhostProvider::alloc(), init] };
let err = Features::from_provider(ProtocolObject::from_ref(&*provider), None, &mut Vec::new())
.unwrap_err();
assert_eq!(err, crate::PredictionError::MissingOutput("ghost".into()));
}