use std::path::{Path, PathBuf};
use objc2::rc::{Retained, autoreleasepool};
use objc2_core_ml::{MLDictionaryFeatureProvider, MLModel, MLModelConfiguration};
use objc2_foundation::NSURL;
use crate::{
CompileError, ComputeUnits, DataType, Features, LoadError, MultiArray, NsErrorInfo,
PredictionError,
};
fn file_url(path: &Path, is_directory: bool) -> Retained<NSURL> {
use std::os::unix::ffi::OsStrExt;
let bytes = std::ffi::CString::new(path.as_os_str().as_bytes())
.expect("callers verify the path exists, so it contains no interior NUL");
unsafe {
NSURL::fileURLWithFileSystemRepresentation_isDirectory_relativeToURL(
core::ptr::NonNull::new(bytes.as_ptr().cast_mut()).expect("CString pointer is non-null"),
is_directory,
None,
)
}
}
#[derive(Debug)]
pub struct Model {
inner: Retained<MLModel>,
description: ModelDescription,
}
unsafe impl Send for Model {}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct AxisRange {
min: usize,
count: usize,
}
impl AxisRange {
#[inline(always)]
pub const fn new(min: usize, count: usize) -> Self {
Self { min, count }
}
#[inline(always)]
pub const fn inclusive(min: usize, max: usize) -> Self {
Self::new(min, max.saturating_sub(min).saturating_add(1))
}
#[inline(always)]
pub const fn min(&self) -> usize {
self.min
}
#[inline(always)]
pub const fn count(&self) -> usize {
self.count
}
}
impl core::fmt::Display for AxisRange {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self.count {
0 => write!(f, "(no size)"),
1 => write!(f, "{}", self.min),
_ => write!(f, "{}..={}", self.min, self.min + self.count - 1),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display)]
#[non_exhaustive]
pub enum UnmeasuredEnumeration {
#[display("no enumerated shape")]
NoShapes,
#[display("sole enumerated shape is not the declared shape")]
SoleShapeIsNotDeclared,
#[display("per-axis ranges do not pin the declared shape")]
SpansDoNotPinDeclaredShape,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display)]
#[non_exhaustive]
pub enum ShapeConstraint {
#[display("fixed")]
Fixed,
#[display("enumerated")]
Enumerated,
#[display("range")]
Range,
#[display("unspecified")]
Unspecified,
#[display("unmeasured({_0})")]
Unmeasured(UnmeasuredEnumeration),
#[display("unknown({_0})")]
Unknown(isize),
}
const RAW_UNSPECIFIED: isize = 1;
const RAW_ENUMERATED: isize = 2;
const RAW_RANGE: isize = 3;
fn classify_shape_constraint(
raw_type: isize,
declared_shape: &[usize],
enumerated_shapes: &[Vec<usize>],
axis_ranges: &[AxisRange],
) -> ShapeConstraint {
match raw_type {
RAW_RANGE => ShapeConstraint::Range,
RAW_UNSPECIFIED => ShapeConstraint::Unspecified,
RAW_ENUMERATED => {
if enumerated_shapes.len() >= 2 {
return ShapeConstraint::Enumerated;
}
if enumerated_shapes.is_empty() {
return ShapeConstraint::Unmeasured(UnmeasuredEnumeration::NoShapes);
}
if enumerated_shapes[0] != declared_shape {
return ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SoleShapeIsNotDeclared);
}
if declared_shape.is_empty()
|| axis_ranges.len() != declared_shape.len()
|| !axis_ranges
.iter()
.zip(declared_shape)
.all(|(range, size)| *range == AxisRange::new(*size, 1))
{
return ShapeConstraint::Unmeasured(UnmeasuredEnumeration::SpansDoNotPinDeclaredShape);
}
ShapeConstraint::Fixed
}
other => ShapeConstraint::Unknown(other),
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct RawShapeConstraint {
raw_type: isize,
enumerated_shapes: Vec<Vec<usize>>,
axis_ranges: Vec<AxisRange>,
}
impl RawShapeConstraint {
pub(crate) const fn new(
raw_type: isize,
enumerated_shapes: Vec<Vec<usize>>,
axis_ranges: Vec<AxisRange>,
) -> Self {
Self {
raw_type,
enumerated_shapes,
axis_ranges,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct FeatureInfo {
name: String,
shape: Vec<usize>,
data_type: Option<DataType>,
optional: bool,
axis_ranges: Vec<AxisRange>,
enumerated_shapes: Vec<Vec<usize>>,
shape_constraint: Option<ShapeConstraint>,
}
impl FeatureInfo {
pub(crate) fn from_parts(
name: String,
shape: Vec<usize>,
data_type: Option<DataType>,
optional: bool,
constraint: Option<RawShapeConstraint>,
) -> Self {
let (axis_ranges, enumerated_shapes, shape_constraint) = match constraint {
None => (Vec::new(), Vec::new(), None),
Some(raw) => {
let verdict = classify_shape_constraint(
raw.raw_type,
&shape,
&raw.enumerated_shapes,
&raw.axis_ranges,
);
(raw.axis_ranges, raw.enumerated_shapes, Some(verdict))
}
};
Self {
name,
shape,
data_type,
optional,
axis_ranges,
enumerated_shapes,
shape_constraint,
}
}
#[inline(always)]
pub fn name(&self) -> &str {
&self.name
}
#[inline(always)]
pub fn shape(&self) -> &[usize] {
&self.shape
}
#[inline(always)]
pub fn axis_ranges(&self) -> &[AxisRange] {
&self.axis_ranges
}
#[inline(always)]
pub fn enumerated_shapes(&self) -> &[Vec<usize>] {
&self.enumerated_shapes
}
#[inline(always)]
pub const fn data_type(&self) -> Option<DataType> {
self.data_type
}
#[inline(always)]
pub const fn is_optional(&self) -> bool {
self.optional
}
#[inline(always)]
pub const fn shape_constraint(&self) -> Option<ShapeConstraint> {
self.shape_constraint
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelDescription {
inputs: Vec<FeatureInfo>,
outputs: Vec<FeatureInfo>,
states: Vec<FeatureInfo>,
}
impl ModelDescription {
pub(crate) const fn from_parts(
inputs: Vec<FeatureInfo>,
outputs: Vec<FeatureInfo>,
states: Vec<FeatureInfo>,
) -> Self {
Self {
inputs,
outputs,
states,
}
}
#[inline(always)]
pub fn inputs(&self) -> &[FeatureInfo] {
&self.inputs
}
#[inline(always)]
pub fn outputs(&self) -> &[FeatureInfo] {
&self.outputs
}
#[inline(always)]
pub fn states(&self) -> &[FeatureInfo] {
&self.states
}
pub fn input(&self, name: &str) -> Option<&FeatureInfo> {
self.inputs.iter().find(|f| f.name == name)
}
pub fn output(&self, name: &str) -> Option<&FeatureInfo> {
self.outputs.iter().find(|f| f.name == name)
}
}
fn snapshot_features(
descriptions: &objc2_foundation::NSDictionary<
objc2_foundation::NSString,
objc2_core_ml::MLFeatureDescription,
>,
) -> Vec<FeatureInfo> {
let mut features = Vec::new();
for name in descriptions.keys() {
let description = descriptions.objectForKey(&name).expect("key from keys()");
let (shape, data_type, raw_constraint) = unsafe {
description
.multiArrayConstraint()
.map_or((Vec::new(), None, None), |constraint| {
let shape_constraint = constraint.shapeConstraint();
let axis_ranges: Vec<AxisRange> = shape_constraint
.sizeRangeForDimension()
.iter()
.map(|range| {
let range = range.rangeValue();
AxisRange::new(range.location, range.length)
})
.collect();
let enumerated_shapes: Vec<Vec<usize>> = shape_constraint
.enumeratedShapes()
.iter()
.map(|dims| dims.iter().map(|n| n.as_usize()).collect())
.collect();
(
constraint.shape().iter().map(|n| n.as_usize()).collect(),
Some(DataType::from_raw(constraint.dataType().0)),
Some(RawShapeConstraint::new(
shape_constraint.r#type().0,
enumerated_shapes,
axis_ranges,
)),
)
})
};
let optional = unsafe { description.isOptional() };
features.push(FeatureInfo::from_parts(
name.to_string(),
shape,
data_type,
optional,
raw_constraint,
));
}
features.sort_by(|a, b| a.name.cmp(&b.name));
features
}
fn snapshot_states(description: &objc2_core_ml::MLModelDescription) -> Vec<FeatureInfo> {
use objc2::runtime::NSObjectProtocol;
if !description.respondsToSelector(objc2::sel!(stateDescriptionsByName)) {
return Vec::new();
}
let states = unsafe { description.stateDescriptionsByName() };
snapshot_features(&states)
}
impl Model {
pub fn load(path: impl AsRef<Path>, units: ComputeUnits) -> Result<Self, LoadError> {
let path = path.as_ref();
if !path.exists() {
return Err(LoadError::NotFound(path.to_path_buf()));
}
let url = file_url(path, path.is_dir());
let configuration = unsafe {
let configuration = MLModelConfiguration::new();
configuration.setComputeUnits(units.to_raw());
configuration
};
let inner =
unsafe { MLModel::modelWithContentsOfURL_configuration_error(&url, &configuration) }
.map_err(|e| LoadError::Native(NsErrorInfo::from_ns_error(&e)))?;
let raw_description = unsafe { inner.modelDescription() };
let (inputs, outputs) = unsafe {
(
snapshot_features(&raw_description.inputDescriptionsByName()),
snapshot_features(&raw_description.outputDescriptionsByName()),
)
};
let states = snapshot_states(&raw_description);
Ok(Self {
inner,
description: ModelDescription::from_parts(inputs, outputs, states),
})
}
#[inline(always)]
pub const fn description(&self) -> &ModelDescription {
&self.description
}
pub(crate) fn raw(&self) -> &MLModel {
&self.inner
}
pub fn predict(&self, inputs: &Features) -> Result<Features, PredictionError> {
autoreleasepool(|_| {
let provider = inputs.to_provider()?;
self.predict_from_provider(&provider, None, inputs.byte_ranges())
})
}
pub fn predict_with(&self, inputs: &[(&str, &MultiArray)]) -> Result<Features, PredictionError> {
autoreleasepool(|_| {
let provider = crate::features::provider_from_pairs(inputs.iter().copied())?;
let known_regions = inputs.iter().map(|(_, a)| a.byte_range()).collect();
self.predict_from_provider(&provider, None, known_regions)
})
}
pub fn predict_with_outputs(
&self,
inputs: &[(&str, &MultiArray)],
outputs: &[&str],
) -> Result<Features, PredictionError> {
autoreleasepool(|_| {
let provider = crate::features::provider_from_pairs(inputs.iter().copied())?;
let known_regions = inputs.iter().map(|(_, a)| a.byte_range()).collect();
self.predict_from_provider(&provider, Some(outputs), known_regions)
})
}
fn predict_from_provider(
&self,
provider: &MLDictionaryFeatureProvider,
wanted: Option<&[&str]>,
mut known_regions: Vec<(usize, usize)>,
) -> Result<Features, PredictionError> {
let outputs = unsafe {
self
.raw()
.predictionFromFeatures_error(objc2::runtime::ProtocolObject::from_ref(provider))
}
.map_err(|e| PredictionError::Native(NsErrorInfo::from_ns_error(&e)))?;
Features::from_provider(&outputs, wanted, &mut known_regions)
}
pub fn compile(source: impl AsRef<Path>) -> Result<PathBuf, CompileError> {
let source = source.as_ref();
if !source.exists() {
return Err(CompileError::NotFound(source.to_path_buf()));
}
let url = file_url(source, source.is_dir());
#[allow(deprecated)]
let compiled = unsafe { MLModel::compileModelAtURL_error(&url) }
.map_err(|e| CompileError::Native(NsErrorInfo::from_ns_error(&e)))?;
let path = compiled.path().expect("compiled model URL has a path");
Ok(PathBuf::from(path.to_string()))
}
pub fn prewarm(path: impl AsRef<Path>, units: ComputeUnits) -> Result<(), LoadError> {
Self::load(path, units).map(drop)
}
pub fn supports_state(&self) -> bool {
use objc2::runtime::NSObjectProtocol;
self.inner.respondsToSelector(objc2::sel!(newState))
}
pub fn make_state(&self) -> Result<crate::State, PredictionError> {
if !self.supports_state() {
return Err(PredictionError::StateUnsupported);
}
Ok(crate::State::from_raw(unsafe { self.inner.newState() }))
}
pub fn predict_with_state(
&self,
inputs: &Features,
state: &mut crate::State,
) -> Result<Features, PredictionError> {
if !self.supports_state() {
return Err(PredictionError::StateUnsupported);
}
autoreleasepool(|_| {
let provider = inputs.to_provider()?;
let outputs = unsafe {
self.inner.predictionFromFeatures_usingState_error(
objc2::runtime::ProtocolObject::from_ref(&*provider),
state.raw(),
)
}
.map_err(|e| PredictionError::Native(NsErrorInfo::from_ns_error(&e)))?;
let mut known_regions = inputs.byte_ranges();
Features::from_provider(&outputs, None, &mut known_regions)
})
}
}
#[cfg_attr(
not(any(
feature = "identity",
feature = "face",
feature = "speaker",
feature = "whisper",
feature = "lid",
test
)),
allow(dead_code, reason = "no door in this feature set holds a `Checked`")
)]
pub(crate) mod contract;
#[cfg(test)]
mod tests;