use std::path::{Path, PathBuf};
use crate::ComputeUnits;
use serde_json::Value;
use crate::audio::whisper::{
error::{ModelError, ModelName},
path_component::single_path_component,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant)]
#[display("{}", self.as_str())]
#[non_exhaustive]
pub enum ModelState {
Unloaded,
Downloading,
Downloaded,
Prewarming,
Prewarmed,
Loading,
Loaded,
Unloading,
}
impl ModelState {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Unloaded => "unloaded",
Self::Downloading => "downloading",
Self::Downloaded => "downloaded",
Self::Prewarming => "prewarming",
Self::Prewarmed => "prewarmed",
Self::Loading => "loading",
Self::Loaded => "loaded",
Self::Unloading => "unloading",
}
}
#[inline(always)]
pub const fn is_busy(&self) -> bool {
matches!(
self,
Self::Downloading | Self::Prewarming | Self::Loading | Self::Unloading
)
}
}
pub type StateCallback = Box<dyn Fn(Option<ModelState>, ModelState) + Send>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, derive_more::Display, derive_more::IsVariant)]
#[display("{}", self.as_str())]
#[non_exhaustive]
pub enum ModelVariant {
Tiny,
TinyEn,
Base,
BaseEn,
Small,
SmallEn,
Medium,
MediumEn,
Large,
LargeV2,
LargeV3,
}
impl ModelVariant {
#[inline(always)]
pub const fn as_str(&self) -> &'static str {
match self {
Self::Tiny => "tiny",
Self::TinyEn => "tiny.en",
Self::Base => "base",
Self::BaseEn => "base.en",
Self::Small => "small",
Self::SmallEn => "small.en",
Self::Medium => "medium",
Self::MediumEn => "medium.en",
Self::Large => "large",
Self::LargeV2 => "large-v2",
Self::LargeV3 => "large-v3",
}
}
#[inline(always)]
pub const fn is_multilingual(&self) -> bool {
!matches!(
self,
Self::TinyEn | Self::BaseEn | Self::SmallEn | Self::MediumEn
)
}
}
#[inline(always)]
pub const fn is_model_multilingual(logits_dim: usize) -> bool {
logits_dim != 51_864
}
pub const fn detect_variant(logits_dim: usize, encoder_dim: usize) -> Option<ModelVariant> {
match logits_dim {
51_865 => match encoder_dim {
384 => Some(ModelVariant::Tiny),
512 => Some(ModelVariant::Base),
768 => Some(ModelVariant::Small),
1024 => Some(ModelVariant::Medium),
1280 => Some(ModelVariant::LargeV2),
_ => None,
},
51_864 => match encoder_dim {
384 => Some(ModelVariant::TinyEn),
512 => Some(ModelVariant::BaseEn),
768 => Some(ModelVariant::SmallEn),
1024 => Some(ModelVariant::MediumEn),
_ => None,
},
51_866 => Some(ModelVariant::LargeV3),
_ => None,
}
}
pub fn detect_model_url(folder: &Path, name: &str, recursive: bool) -> Result<PathBuf, ModelError> {
let name = single_path_component(name)
.map_err(|defect| ModelError::ModelName(ModelName::new(name.to_owned(), defect.reason())))?;
let compiled = folder.join(format!("{name}.mlmodelc"));
if recursive {
if compiled.exists() {
return Ok(compiled);
}
let target = format!("{name}.mlmodelc");
if let Some(found) = find_named_recursive(folder, &target) {
return Ok(found);
}
return Err(ModelError::NotFound(vec![compiled]));
}
if compiled.exists() {
return Ok(compiled);
}
let package = folder
.join(format!("{name}.mlpackage"))
.join("Data/com.apple.CoreML/model.mlmodel");
if package.exists() {
return Ok(package);
}
Err(ModelError::NotFound(vec![compiled, package]))
}
fn find_named_recursive(folder: &Path, target_name: &str) -> Option<PathBuf> {
let mut worklist = vec![folder.to_path_buf()];
while let Some(dir) = worklist.pop() {
let Ok(entries) = std::fs::read_dir(&dir) else {
continue;
};
for entry in entries.flatten() {
let path = entry.path();
if path.file_name().and_then(|n| n.to_str()) == Some(target_name) {
return Some(path);
}
if entry.file_type().is_ok_and(|t| t.is_dir()) {
worklist.push(path);
}
}
}
None
}
pub fn glob_match(pattern: &str, name: &str) -> bool {
let pat: Vec<char> = pattern.chars().collect();
let text: Vec<char> = name.chars().collect();
let (mut pi, mut ti) = (0usize, 0usize);
let mut star: Option<(usize, usize)> = None;
while ti < text.len() {
if pi < pat.len() && (pat[pi] == '?' || pat[pi] == text[ti]) {
pi += 1;
ti += 1;
} else if pi < pat.len() && pat[pi] == '*' {
star = Some((pi + 1, ti));
pi += 1;
} else if let Some((resume_pi, matched_ti)) = star {
ti = matched_ti + 1;
pi = resume_pi;
star = Some((resume_pi, ti));
} else {
return false;
}
}
while pat.get(pi) == Some(&'*') {
pi += 1;
}
pi == pat.len()
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelInfo {
name: String,
version: Option<String>,
variant: Option<String>,
compute: ComputeUnits,
}
impl ModelInfo {
pub fn try_new(
name: impl Into<String>,
version: Option<String>,
variant: Option<String>,
compute: ComputeUnits,
) -> Result<Self, ModelError> {
let name = name.into();
if name.is_empty() {
return Err(ModelError::EmptyName);
}
Ok(Self {
name,
version,
variant,
compute,
})
}
#[inline(always)]
pub fn name(&self) -> &str {
self.name.as_str()
}
#[inline(always)]
pub fn version(&self) -> Option<&str> {
self.version.as_deref()
}
#[inline(always)]
pub fn variant(&self) -> Option<&str> {
self.variant.as_deref()
}
#[inline(always)]
pub const fn compute(&self) -> ComputeUnits {
self.compute
}
pub fn download_pattern(&self) -> String {
format!(
"{}/{}/{}/*",
self.name,
self.version.as_deref().unwrap_or("*"),
self.variant.as_deref().unwrap_or("*"),
)
}
pub fn find_base_folder(&self, path: &Path) -> Option<PathBuf> {
let mut current = path.to_path_buf();
loop {
if current.components().count() <= 1 {
return None;
}
if current.file_name().and_then(|n| n.to_str()) == Some(self.name.as_str()) {
current.pop();
return Some(current);
}
if !current.pop() {
return None;
}
}
}
}
pub const DEFAULT_FALLBACK_MODEL_NAME: &str = "openai_whisper-base";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelSupport {
default_model: String,
supported: Vec<String>,
disabled: Vec<String>,
}
impl ModelSupport {
pub fn new(
default_model: impl Into<String>,
supported: impl Into<Vec<String>>,
disabled: impl Into<Vec<String>>,
) -> Self {
Self {
default_model: default_model.into(),
supported: supported.into(),
disabled: disabled.into(),
}
}
#[inline(always)]
pub fn default_model(&self) -> &str {
self.default_model.as_str()
}
#[inline(always)]
pub const fn supported_slice(&self) -> &[String] {
self.supported.as_slice()
}
#[inline(always)]
pub const fn disabled_slice(&self) -> &[String] {
self.disabled.as_slice()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DeviceSupport {
chips: Option<String>,
identifiers: Vec<String>,
models: ModelSupport,
}
impl DeviceSupport {
pub fn new(
chips: Option<String>,
identifiers: impl Into<Vec<String>>,
models: ModelSupport,
) -> Self {
Self {
chips,
identifiers: identifiers.into(),
models,
}
}
#[inline(always)]
pub fn chips(&self) -> Option<&str> {
self.chips.as_deref()
}
#[inline(always)]
pub const fn identifiers_slice(&self) -> &[String] {
self.identifiers.as_slice()
}
#[inline(always)]
pub const fn models(&self) -> &ModelSupport {
&self.models
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SupportConfig {
device_supports: Vec<DeviceSupport>,
known_models: Vec<String>,
default_support: DeviceSupport,
}
impl SupportConfig {
fn from_device_supports(device_supports: Vec<DeviceSupport>) -> Self {
let mut known_models: Vec<String> = Vec::new();
for ds in &device_supports {
for model in ds.models().supported_slice() {
if !known_models.contains(model) {
known_models.push(model.clone());
}
}
}
let device_supports: Vec<DeviceSupport> = device_supports
.into_iter()
.map(|ds| {
let disabled: Vec<String> = known_models
.iter()
.filter(|m| !ds.models().supported_slice().contains(m))
.cloned()
.collect();
DeviceSupport::new(
ds.chips().map(str::to_string),
ds.identifiers_slice().to_vec(),
ModelSupport::new(
ds.models().default_model().to_string(),
ds.models().supported_slice().to_vec(),
disabled,
),
)
})
.collect();
let default_support = DeviceSupport::new(
None,
Vec::new(),
ModelSupport::new(
DEFAULT_FALLBACK_MODEL_NAME,
known_models.clone(),
Vec::new(),
),
);
Self {
device_supports,
known_models,
default_support,
}
}
pub fn from_json(json: &str) -> Result<Self, ModelError> {
let value: Value =
serde_json::from_str(json).map_err(|e| ModelError::InvalidSupportConfig(e.to_string()))?;
let entries = value
.get("device_support")
.and_then(Value::as_array)
.ok_or_else(|| {
ModelError::InvalidSupportConfig("missing `device_support` array".to_string())
})?;
let device_supports = entries
.iter()
.map(device_support_from_json)
.collect::<Result<Vec<_>, _>>()?;
Ok(Self::from_device_supports(device_supports))
}
pub fn fallback() -> Self {
Self::from_device_supports(fallback_device_supports())
}
pub fn support_for(&self, device_identifier: &str) -> ModelSupport {
let mut best: Option<(&DeviceSupport, usize)> = None;
for ds in &self.device_supports {
for identifier in ds.identifiers_slice() {
if !device_identifier.starts_with(identifier.as_str()) {
continue;
}
let len = identifier.len();
let better = match best {
None => true,
Some((_, best_len)) => len > best_len,
};
if better {
best = Some((ds, len));
}
}
}
best.map_or_else(
|| self.default_support.models().clone(),
|(ds, _)| ds.models().clone(),
)
}
#[inline(always)]
pub const fn device_supports_slice(&self) -> &[DeviceSupport] {
self.device_supports.as_slice()
}
#[inline(always)]
pub const fn known_models_slice(&self) -> &[String] {
self.known_models.as_slice()
}
#[inline(always)]
pub const fn default_support(&self) -> &DeviceSupport {
&self.default_support
}
}
fn device_support_from_json(value: &Value) -> Result<DeviceSupport, ModelError> {
let chips = value
.get("chips")
.and_then(Value::as_str)
.map(str::to_string);
let identifiers = string_array(value, "identifiers")?;
let models_value = value.get("models").ok_or_else(|| {
ModelError::InvalidSupportConfig("device support entry missing `models`".to_string())
})?;
let models = model_support_from_json(models_value)?;
Ok(DeviceSupport::new(chips, identifiers, models))
}
fn model_support_from_json(value: &Value) -> Result<ModelSupport, ModelError> {
let default_model = value
.get("default")
.and_then(Value::as_str)
.ok_or_else(|| ModelError::InvalidSupportConfig("models entry missing `default`".to_string()))?
.to_string();
let supported = string_array(value, "supported")?;
Ok(ModelSupport::new(default_model, supported, Vec::new()))
}
fn string_array(value: &Value, key: &str) -> Result<Vec<String>, ModelError> {
value
.get(key)
.and_then(Value::as_array)
.ok_or_else(|| ModelError::InvalidSupportConfig(format!("missing `{key}` array")))?
.iter()
.map(|entry| {
entry
.as_str()
.map(str::to_string)
.ok_or_else(|| ModelError::InvalidSupportConfig(format!("`{key}` entry is not a string")))
})
.collect()
}
fn strs(items: &[&str]) -> Vec<String> {
items.iter().map(|s| (*s).to_string()).collect()
}
fn fallback_device_supports() -> Vec<DeviceSupport> {
vec![
DeviceSupport::new(
Some("A12, A13, S9, S10".to_string()),
strs(&[
"iPhone11", "iPhone12", "iPad12,1", "iPad12,2", "Watch7", "Watch8",
]),
ModelSupport::new(
"openai_whisper-tiny",
strs(&[
"openai_whisper-base",
"openai_whisper-base.en",
"openai_whisper-tiny",
"openai_whisper-tiny.en",
]),
Vec::new(),
),
),
DeviceSupport::new(
Some("A16, A17 Pro, A18".to_string()),
strs(&[
"iPhone15", "iPhone16", "iPhone17", "iPad15,7", "iPad15,8", "iPad16,1", "iPad16,2",
]),
ModelSupport::new(
"openai_whisper-base",
strs(&[
"openai_whisper-tiny",
"openai_whisper-tiny.en",
"openai_whisper-base",
"openai_whisper-base.en",
"openai_whisper-small",
"openai_whisper-small.en",
"openai_whisper-large-v2_949MB",
"openai_whisper-large-v2_turbo_955MB",
"openai_whisper-large-v3_947MB",
"openai_whisper-large-v3_turbo_954MB",
"distil-whisper_distil-large-v3_594MB",
"distil-whisper_distil-large-v3_turbo_600MB",
"openai_whisper-large-v3-v20240930_626MB",
"openai_whisper-large-v3-v20240930_turbo_632MB",
]),
Vec::new(),
),
),
DeviceSupport::new(
Some("M1".to_string()),
strs(&[
"MacBookPro17,1",
"MacBookPro18,1",
"MacBookPro18,2",
"MacBookPro18,3",
"MacBookPro18,4",
"MacBookAir10,1",
"Macmini9,1",
"iMac21,1",
"iMac21,2",
"Mac13",
"iPad13,4",
"iPad13,5",
"iPad13,6",
"iPad13,7",
"iPad13,8",
"iPad13,9",
"iPad13,10",
"iPad13,11",
"iPad13,16",
"iPad13,17",
]),
ModelSupport::new(
"openai_whisper-large-v3-v20240930_626MB",
strs(&[
"openai_whisper-tiny",
"openai_whisper-tiny.en",
"openai_whisper-base",
"openai_whisper-base.en",
"openai_whisper-small",
"openai_whisper-small.en",
"openai_whisper-large-v2",
"openai_whisper-large-v2_949MB",
"openai_whisper-large-v3",
"openai_whisper-large-v3_947MB",
"distil-whisper_distil-large-v3",
"distil-whisper_distil-large-v3_594MB",
"openai_whisper-large-v3-v20240930_626MB",
]),
Vec::new(),
),
),
DeviceSupport::new(
Some("M2, M3, M4".to_string()),
strs(&[
"Mac14",
"Mac15",
"Mac16",
"iPad14,3",
"iPad14,4",
"iPad14,5",
"iPad14,6",
"iPad14,8",
"iPad14,9",
"iPad14,10",
"iPad14,11",
"iPad15",
"iPad16",
]),
ModelSupport::new(
"openai_whisper-large-v3-v20240930",
strs(&[
"openai_whisper-tiny",
"openai_whisper-tiny.en",
"openai_whisper-base",
"openai_whisper-base.en",
"openai_whisper-small",
"openai_whisper-small.en",
"openai_whisper-large-v2",
"openai_whisper-large-v2_949MB",
"openai_whisper-large-v2_turbo",
"openai_whisper-large-v2_turbo_955MB",
"openai_whisper-large-v3",
"openai_whisper-large-v3_947MB",
"openai_whisper-large-v3_turbo",
"openai_whisper-large-v3_turbo_954MB",
"distil-whisper_distil-large-v3",
"distil-whisper_distil-large-v3_594MB",
"distil-whisper_distil-large-v3_turbo",
"distil-whisper_distil-large-v3_turbo_600MB",
"openai_whisper-large-v3-v20240930",
"openai_whisper-large-v3-v20240930_turbo",
"openai_whisper-large-v3-v20240930_626MB",
"openai_whisper-large-v3-v20240930_turbo_632MB",
]),
Vec::new(),
),
),
DeviceSupport::new(
Some("A14".to_string()),
strs(&["iPhone13", "iPad13,1", "iPad13,2", "iPad13,18", "iPad13,19"]),
ModelSupport::new(
"openai_whisper-base",
strs(&[
"openai_whisper-tiny",
"openai_whisper-tiny.en",
"openai_whisper-base",
"openai_whisper-base.en",
"openai_whisper-small",
"openai_whisper-small.en",
]),
Vec::new(),
),
),
DeviceSupport::new(
Some("A15".to_string()),
strs(&["iPhone14", "iPad14,1", "iPad14,2"]),
ModelSupport::new(
"openai_whisper-base",
strs(&[
"openai_whisper-tiny",
"openai_whisper-tiny.en",
"openai_whisper-base",
"openai_whisper-base.en",
"openai_whisper-small",
"openai_whisper-small.en",
"openai_whisper-large-v2_949MB",
"openai_whisper-large-v2_turbo_955MB",
"openai_whisper-large-v3_947MB",
"openai_whisper-large-v3_turbo_954MB",
"distil-whisper_distil-large-v3_594MB",
"distil-whisper_distil-large-v3_turbo_600MB",
"openai_whisper-large-v3-v20240930_626MB",
"openai_whisper-large-v3-v20240930_turbo_632MB",
]),
Vec::new(),
),
),
]
}
pub fn device_identifier() -> String {
let name = c"hw.model";
let mut size: usize = 0;
let query = unsafe {
libc::sysctlbyname(
name.as_ptr(),
std::ptr::null_mut(),
&mut size,
std::ptr::null_mut(),
0,
)
};
if query != 0 || size == 0 {
return "unknown".to_string();
}
let mut buf = vec![0_u8; size];
let fill = unsafe {
libc::sysctlbyname(
name.as_ptr(),
buf.as_mut_ptr().cast(),
&mut size,
std::ptr::null_mut(),
0,
)
};
if fill != 0 {
return "unknown".to_string();
}
let end = buf.iter().position(|&b| b == 0).unwrap_or(buf.len());
String::from_utf8_lossy(&buf[..end]).into_owned()
}
pub trait ModelLoader {
fn resolve(&self, folder: &Path) -> Result<ResolvedModels, ModelError>;
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResolvedModels {
mel: PathBuf,
encoder: PathBuf,
decoder: PathBuf,
}
impl ResolvedModels {
pub fn new(
mel: impl Into<PathBuf>,
encoder: impl Into<PathBuf>,
decoder: impl Into<PathBuf>,
) -> Self {
Self {
mel: mel.into(),
encoder: encoder.into(),
decoder: decoder.into(),
}
}
#[inline(always)]
pub fn mel_ref(&self) -> &Path {
self.mel.as_path()
}
#[inline(always)]
pub fn encoder_ref(&self) -> &Path {
self.encoder.as_path()
}
#[inline(always)]
pub fn decoder_ref(&self) -> &Path {
self.decoder.as_path()
}
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct LocalModelLoader;
impl LocalModelLoader {
#[inline(always)]
pub const fn new() -> Self {
Self
}
}
impl ModelLoader for LocalModelLoader {
fn resolve(&self, folder: &Path) -> Result<ResolvedModels, ModelError> {
let mel = detect_model_url(folder, "MelSpectrogram", false)?;
let encoder = detect_model_url(folder, "AudioEncoder", false)?;
let decoder = detect_model_url(folder, "TextDecoder", false)?;
Ok(ResolvedModels::new(mel, encoder, decoder))
}
}
pub mod manager;
#[cfg(test)]
mod tests;