use std::collections::BTreeMap;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum TaskKind {
Vad,
Asr,
Diar,
SourceSeparation,
AudioGeneration,
Tts,
VoiceCloning,
VoiceConversion,
SpeechToSpeech,
Alignment,
VoiceDesign,
SpeakerRecognition,
Svc,
Midi,
}
impl TaskKind {
pub fn as_str(self) -> &'static str {
match self {
TaskKind::Vad => "vad",
TaskKind::Asr => "asr",
TaskKind::Diar => "diar",
TaskKind::SourceSeparation => "sep",
TaskKind::AudioGeneration => "gen",
TaskKind::Tts => "tts",
TaskKind::VoiceCloning => "clon",
TaskKind::VoiceConversion => "vc",
TaskKind::SpeechToSpeech => "s2s",
TaskKind::Alignment => "align",
TaskKind::VoiceDesign => "vdes",
TaskKind::SpeakerRecognition => "spk",
TaskKind::Svc => "svc",
TaskKind::Midi => "midi",
}
}
}
impl Serialize for TaskKind {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_str(self.as_str())
}
}
impl<'de> Deserialize<'de> for TaskKind {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let s = String::deserialize(deserializer)?;
match s.as_str() {
"vad" => Ok(TaskKind::Vad),
"asr" => Ok(TaskKind::Asr),
"diar" => Ok(TaskKind::Diar),
"sep" => Ok(TaskKind::SourceSeparation),
"gen" => Ok(TaskKind::AudioGeneration),
"tts" => Ok(TaskKind::Tts),
"clon" => Ok(TaskKind::VoiceCloning),
"vc" => Ok(TaskKind::VoiceConversion),
"s2s" => Ok(TaskKind::SpeechToSpeech),
"align" => Ok(TaskKind::Alignment),
"vdes" => Ok(TaskKind::VoiceDesign),
"spk" => Ok(TaskKind::SpeakerRecognition),
"svc" => Ok(TaskKind::Svc),
"midi" => Ok(TaskKind::Midi),
other => Err(serde::de::Error::unknown_variant(
other,
&[
"vad", "asr", "diar", "sep", "gen", "tts", "clon", "vc", "s2s", "align",
"vdes", "spk", "svc", "midi",
],
)),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum RunMode {
Offline,
Streaming,
}
impl RunMode {
pub fn as_str(self) -> &'static str {
match self {
RunMode::Offline => "offline",
RunMode::Streaming => "streaming",
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum ModelFamily {
SileroVad,
MarblenetVad,
Qwen3Asr,
CitrinetAsr,
SenseAsr,
FunAsrNano,
HiggsAudioStt,
HviskeAsr,
KrokoAsr,
NemotronAsr,
ParakeetTdt,
VibevoiceAsr,
VibevoiceAsrStreaming,
FireredAudio,
Granite5Asr,
Audio8Asr,
Audio8Tts,
Vibeasr,
MoonshineAsr,
NiagaraAsr,
Qwen3Tts,
Confucius4Tts,
DotsTts,
FishAudio,
GlmTts,
HiggsAudioTts,
IndexTts2,
IrodoriTts,
MossTtsLocal,
MossTtsNano,
MossVoicegen,
Neutts,
Outetts,
PocketTts,
VietneuTts,
Fireredtts3,
F5Tts,
MagpieTts,
Personaplex,
MinimaxH3,
Sanotts,
SoproTts,
MiraTts,
Cosyvoice3,
BreezeTts,
SopranoTts,
EchoTts,
Voxcpm1,
KokoroTts,
SortformerDiar,
SortformerDiarV2,
SeedVc,
Rvc,
Meanvc2,
Chatterbox,
ChatterboxTurbo,
Vevo2,
Voxcpm2,
AceStep,
Htdemucs,
MelBandRoformer,
BsRoformer,
Muscriptor,
Omnivoice,
StableAudio,
MinimaxMusic3,
Supertonic,
VoxtralRealtime,
Audiosr,
BuiltinAudioUtils,
ControlFoley,
MidashEnglmGen,
Miocodec,
Miotts,
Vibevoice,
Dramabox,
Heartmula,
InflectV2,
Qwen3ForcedAligner,
MmsForcedAligner,
Custom(String),
}
impl ModelFamily {
pub fn from_path(path: &str) -> Option<ModelFamily> {
let lower = path.to_ascii_lowercase();
if let Some((_, f)) = lower.rsplit_once("#family=") {
return Some(ModelFamily::from(f.trim()));
}
const KEYWORDS: &[(&str, ModelFamily)] = &[
("silero_vad", ModelFamily::SileroVad),
("silero-vad", ModelFamily::SileroVad),
("marblenet_vad", ModelFamily::MarblenetVad),
("marblenet-vad", ModelFamily::MarblenetVad),
("marblenet", ModelFamily::MarblenetVad),
("qwen3_asr", ModelFamily::Qwen3Asr),
("qwen3-asr", ModelFamily::Qwen3Asr),
("qwen3_tts", ModelFamily::Qwen3Tts),
("qwen3-tts", ModelFamily::Qwen3Tts),
("qwen3_forced_aligner", ModelFamily::Qwen3ForcedAligner),
("qwen3-forced-aligner", ModelFamily::Qwen3ForcedAligner),
("qwen3", ModelFamily::Qwen3Asr),
("citrinet", ModelFamily::CitrinetAsr),
("sense_asr", ModelFamily::SenseAsr),
("sense-asr", ModelFamily::SenseAsr),
("sensevoice", ModelFamily::SenseAsr),
("fun-asr-nano", ModelFamily::FunAsrNano),
("fun_asr_nano", ModelFamily::FunAsrNano),
("funasr", ModelFamily::FunAsrNano),
("higgs_audio_stt", ModelFamily::HiggsAudioStt),
("higgs_audio_tts", ModelFamily::HiggsAudioTts),
("higgs", ModelFamily::HiggsAudioStt),
("hviske", ModelFamily::HviskeAsr),
("kroko", ModelFamily::KrokoAsr),
("nemotron", ModelFamily::NemotronAsr),
("parakeet", ModelFamily::ParakeetTdt),
(
"vibevoice_asr_streaming",
ModelFamily::VibevoiceAsrStreaming,
),
(
"vibevoice-asr-streaming",
ModelFamily::VibevoiceAsrStreaming,
),
("vibevoice_asr", ModelFamily::VibevoiceAsr),
("firered_audio", ModelFamily::FireredAudio),
("firered-audio", ModelFamily::FireredAudio),
("granite5asr", ModelFamily::Granite5Asr),
("granite-speech", ModelFamily::Granite5Asr),
("granite_speech", ModelFamily::Granite5Asr),
("granite", ModelFamily::Granite5Asr),
("audio8_asr", ModelFamily::Audio8Asr),
("audio8-asr", ModelFamily::Audio8Asr),
("audio8asr", ModelFamily::Audio8Asr),
("arkasr", ModelFamily::Audio8Asr),
("audio8_tts", ModelFamily::Audio8Tts),
("audio8-tts", ModelFamily::Audio8Tts),
("audio8tts", ModelFamily::Audio8Tts),
("fireredtts3", ModelFamily::Fireredtts3),
("firered_tts3", ModelFamily::Fireredtts3),
("firered-tts3", ModelFamily::Fireredtts3),
("vibevoice", ModelFamily::Vibevoice),
("confucius", ModelFamily::Confucius4Tts),
("dots_tts", ModelFamily::DotsTts),
("dots-tts", ModelFamily::DotsTts),
("fish_audio", ModelFamily::FishAudio),
("fish-audio", ModelFamily::FishAudio),
("glm_tts", ModelFamily::GlmTts),
("glm-tts", ModelFamily::GlmTts),
("index_tts2", ModelFamily::IndexTts2),
("index-tts2", ModelFamily::IndexTts2),
("irodori", ModelFamily::IrodoriTts),
("moss-tts-nano", ModelFamily::MossTtsNano),
("moss_tts_nano", ModelFamily::MossTtsNano),
("moss-tts-local", ModelFamily::MossTtsLocal),
("moss_tts_local", ModelFamily::MossTtsLocal),
("moss_voicegen", ModelFamily::MossVoicegen),
("moss-voicegen", ModelFamily::MossVoicegen),
("moss", ModelFamily::MossTtsNano),
("neutts", ModelFamily::Neutts),
("outetts", ModelFamily::Outetts),
("pocket_tts", ModelFamily::PocketTts),
("pocket-tts", ModelFamily::PocketTts),
("vietneu", ModelFamily::VietneuTts),
("f5_tts", ModelFamily::F5Tts),
("f5-tts", ModelFamily::F5Tts),
("magpie", ModelFamily::MagpieTts),
("personaplex", ModelFamily::Personaplex),
("soprano_tts", ModelFamily::SopranoTts),
("soprano-tts", ModelFamily::SopranoTts),
("soprano", ModelFamily::SopranoTts),
("echo_tts", ModelFamily::EchoTts),
("echo-tts", ModelFamily::EchoTts),
("echo", ModelFamily::EchoTts),
("kokoro_tts", ModelFamily::KokoroTts),
("kokoro-tts", ModelFamily::KokoroTts),
("kokoro", ModelFamily::KokoroTts),
("voxcpm1", ModelFamily::Voxcpm1),
("voxcpm-1", ModelFamily::Voxcpm1),
("sanotts", ModelFamily::Sanotts),
("sano-tts", ModelFamily::Sanotts),
("sano_tts", ModelFamily::Sanotts),
("sopro_tts", ModelFamily::SoproTts),
("sopro-tts", ModelFamily::SoproTts),
("sopro", ModelFamily::SoproTts),
("mira_tts", ModelFamily::MiraTts),
("mira-tts", ModelFamily::MiraTts),
("miratts", ModelFamily::MiraTts),
("mira", ModelFamily::MiraTts),
("cosyvoice3", ModelFamily::Cosyvoice3),
("cosyvoice-3", ModelFamily::Cosyvoice3),
("cosyvoice", ModelFamily::Cosyvoice3),
("breeze_tts", ModelFamily::BreezeTts),
("breeze-tts", ModelFamily::BreezeTts),
("breeze", ModelFamily::BreezeTts),
("vibeasr", ModelFamily::Vibeasr),
("vibe-asr", ModelFamily::Vibeasr),
("moonshine_asr", ModelFamily::MoonshineAsr),
("moonshine-asr", ModelFamily::MoonshineAsr),
("moonshine", ModelFamily::MoonshineAsr),
("niagara_asr", ModelFamily::NiagaraAsr),
("niagara-asr", ModelFamily::NiagaraAsr),
("niagara", ModelFamily::NiagaraAsr),
("minimax_h3", ModelFamily::MinimaxH3),
("minimax-h3", ModelFamily::MinimaxH3),
("sortformer_diar_v2", ModelFamily::SortformerDiarV2),
("sortformer-diar-v2", ModelFamily::SortformerDiarV2),
("sortformer_diar", ModelFamily::SortformerDiar),
("sortformer-diar", ModelFamily::SortformerDiar),
("sortformer", ModelFamily::SortformerDiar),
("seed_vc", ModelFamily::SeedVc),
("seed-vc", ModelFamily::SeedVc),
("seedvc", ModelFamily::SeedVc),
("rvc", ModelFamily::Rvc),
("meanvc2", ModelFamily::Meanvc2),
("mean-vc2", ModelFamily::Meanvc2),
("chatterbox_turbo", ModelFamily::ChatterboxTurbo),
("chatterbox-turbo", ModelFamily::ChatterboxTurbo),
("chatterbox", ModelFamily::Chatterbox),
("vevo2", ModelFamily::Vevo2),
("voxcpm2", ModelFamily::Voxcpm2),
("ace_step", ModelFamily::AceStep),
("ace-step", ModelFamily::AceStep),
("htdemucs", ModelFamily::Htdemucs),
("demucs", ModelFamily::Htdemucs),
("mel-band-roformer", ModelFamily::MelBandRoformer),
("mel_band_roformer", ModelFamily::MelBandRoformer),
("bs-roformer", ModelFamily::BsRoformer),
("bs_roformer", ModelFamily::BsRoformer),
("muscriptor", ModelFamily::Muscriptor),
("omnivoice", ModelFamily::Omnivoice),
("stable_audio", ModelFamily::StableAudio),
("stable-audio", ModelFamily::StableAudio),
("minimax_music3", ModelFamily::MinimaxMusic3),
("minimax-music3", ModelFamily::MinimaxMusic3),
("supertonic", ModelFamily::Supertonic),
("voxtral-realtime", ModelFamily::VoxtralRealtime),
("voxtral", ModelFamily::VoxtralRealtime),
("audiosr", ModelFamily::Audiosr),
("audio-sr", ModelFamily::Audiosr),
("controlfoley", ModelFamily::ControlFoley),
("control-foley", ModelFamily::ControlFoley),
("midashenglm_gen", ModelFamily::MidashEnglmGen),
("midashenglm-gen", ModelFamily::MidashEnglmGen),
("midashen", ModelFamily::MidashEnglmGen),
("builtin_audio_utils", ModelFamily::BuiltinAudioUtils),
("builtin-audio-utils", ModelFamily::BuiltinAudioUtils),
("deepfilternet", ModelFamily::BuiltinAudioUtils),
("rnnoise", ModelFamily::BuiltinAudioUtils),
("zipenhancer", ModelFamily::BuiltinAudioUtils),
("gtcrn", ModelFamily::BuiltinAudioUtils),
("flashsr", ModelFamily::BuiltinAudioUtils),
("miocodec", ModelFamily::Miocodec),
("miotts", ModelFamily::Miotts),
("dramabox", ModelFamily::Dramabox),
("heartmula", ModelFamily::Heartmula),
("inflect", ModelFamily::InflectV2),
("mms_forced_aligner", ModelFamily::MmsForcedAligner),
("mms", ModelFamily::MmsForcedAligner),
("forced-aligner", ModelFamily::Qwen3ForcedAligner),
];
KEYWORDS
.iter()
.find(|(k, _)| lower.contains(k))
.map(|(_, family)| family.clone())
}
pub fn as_str(&self) -> &str {
match self {
ModelFamily::SileroVad => "silero_vad",
ModelFamily::MarblenetVad => "marblenet_vad",
ModelFamily::Qwen3Asr => "qwen3_asr",
ModelFamily::CitrinetAsr => "citrinet_asr",
ModelFamily::SenseAsr => "sense_asr",
ModelFamily::FunAsrNano => "fun_asr_nano",
ModelFamily::HiggsAudioStt => "higgs_audio_stt",
ModelFamily::HviskeAsr => "hviske_asr",
ModelFamily::KrokoAsr => "kroko_asr",
ModelFamily::NemotronAsr => "nemotron_asr",
ModelFamily::ParakeetTdt => "parakeet_tdt",
ModelFamily::VibevoiceAsr => "vibevoice_asr",
ModelFamily::VibevoiceAsrStreaming => "vibevoice_asr_streaming",
ModelFamily::FireredAudio => "firered_audio",
ModelFamily::Granite5Asr => "granite5asr",
ModelFamily::Audio8Asr => "audio8_asr",
ModelFamily::Audio8Tts => "audio8_tts",
ModelFamily::Vibeasr => "vibeasr",
ModelFamily::MoonshineAsr => "moonshine_asr",
ModelFamily::NiagaraAsr => "niagara_asr",
ModelFamily::Qwen3Tts => "qwen3_tts",
ModelFamily::Confucius4Tts => "confucius4_tts",
ModelFamily::DotsTts => "dots_tts",
ModelFamily::FishAudio => "fish_audio",
ModelFamily::GlmTts => "glm_tts",
ModelFamily::HiggsAudioTts => "higgs_audio_tts",
ModelFamily::IndexTts2 => "index_tts2",
ModelFamily::IrodoriTts => "irodori_tts",
ModelFamily::MossTtsLocal => "moss_tts_local",
ModelFamily::MossTtsNano => "moss_tts_nano",
ModelFamily::MossVoicegen => "moss_voicegen",
ModelFamily::Neutts => "neutts",
ModelFamily::Outetts => "outetts",
ModelFamily::PocketTts => "pocket_tts",
ModelFamily::VietneuTts => "vietneu_tts",
ModelFamily::Fireredtts3 => "fireredtts3",
ModelFamily::F5Tts => "f5_tts",
ModelFamily::MagpieTts => "magpie_tts",
ModelFamily::Personaplex => "personaplex",
ModelFamily::SopranoTts => "soprano_tts",
ModelFamily::EchoTts => "echo_tts",
ModelFamily::Voxcpm1 => "voxcpm1",
ModelFamily::KokoroTts => "kokoro_tts",
ModelFamily::MinimaxH3 => "minimax_h3",
ModelFamily::Sanotts => "sanotts",
ModelFamily::SoproTts => "sopro_tts",
ModelFamily::MiraTts => "mira_tts",
ModelFamily::Cosyvoice3 => "cosyvoice3",
ModelFamily::BreezeTts => "breeze_tts",
ModelFamily::SortformerDiar => "sortformer_diar",
ModelFamily::SortformerDiarV2 => "sortformer_diar_v2",
ModelFamily::SeedVc => "seed_vc",
ModelFamily::Rvc => "rvc",
ModelFamily::Meanvc2 => "meanvc2",
ModelFamily::Chatterbox => "chatterbox",
ModelFamily::ChatterboxTurbo => "chatterbox_turbo",
ModelFamily::Vevo2 => "vevo2",
ModelFamily::Voxcpm2 => "voxcpm2",
ModelFamily::AceStep => "ace_step",
ModelFamily::Htdemucs => "htdemucs",
ModelFamily::MelBandRoformer => "mel_band_roformer",
ModelFamily::BsRoformer => "bs_roformer",
ModelFamily::Muscriptor => "muscriptor",
ModelFamily::Omnivoice => "omnivoice",
ModelFamily::StableAudio => "stable_audio",
ModelFamily::MinimaxMusic3 => "minimax_music3",
ModelFamily::Supertonic => "supertonic",
ModelFamily::VoxtralRealtime => "voxtral_realtime",
ModelFamily::Audiosr => "audiosr",
ModelFamily::BuiltinAudioUtils => "builtin_audio_utils",
ModelFamily::ControlFoley => "controlfoley",
ModelFamily::MidashEnglmGen => "midashenglm_gen",
ModelFamily::Miocodec => "miocodec",
ModelFamily::Miotts => "miotts",
ModelFamily::Vibevoice => "vibevoice",
ModelFamily::Dramabox => "dramabox",
ModelFamily::Heartmula => "heartmula",
ModelFamily::InflectV2 => "inflect_v2",
ModelFamily::Qwen3ForcedAligner => "qwen3_forced_aligner",
ModelFamily::MmsForcedAligner => "mms_forced_aligner",
ModelFamily::Custom(s) => s,
}
}
}
impl From<&str> for ModelFamily {
fn from(s: &str) -> Self {
match s {
"silero_vad" => ModelFamily::SileroVad,
"marblenet_vad" => ModelFamily::MarblenetVad,
"qwen3_asr" => ModelFamily::Qwen3Asr,
"citrinet_asr" => ModelFamily::CitrinetAsr,
"sense_asr" => ModelFamily::SenseAsr,
"fun_asr_nano" => ModelFamily::FunAsrNano,
"higgs_audio_stt" => ModelFamily::HiggsAudioStt,
"hviske_asr" => ModelFamily::HviskeAsr,
"kroko_asr" => ModelFamily::KrokoAsr,
"nemotron_asr" => ModelFamily::NemotronAsr,
"parakeet_tdt" => ModelFamily::ParakeetTdt,
"vibevoice_asr" => ModelFamily::VibevoiceAsr,
"vibevoice_asr_streaming" => ModelFamily::VibevoiceAsrStreaming,
"firered_audio" => ModelFamily::FireredAudio,
"granite5asr" => ModelFamily::Granite5Asr,
"granite_speech5_asr" => ModelFamily::Granite5Asr,
"granite_speech" => ModelFamily::Granite5Asr,
"granite_speech5_ctc" => ModelFamily::Granite5Asr,
"audio8_asr" => ModelFamily::Audio8Asr,
"arkasr" => ModelFamily::Audio8Asr,
"audio8_tts" => ModelFamily::Audio8Tts,
"vibeasr" => ModelFamily::Vibeasr,
"moonshine_asr" => ModelFamily::MoonshineAsr,
"niagara_asr" => ModelFamily::NiagaraAsr,
"qwen3_tts" => ModelFamily::Qwen3Tts,
"confucius4_tts" => ModelFamily::Confucius4Tts,
"dots_tts" => ModelFamily::DotsTts,
"fish_audio" => ModelFamily::FishAudio,
"glm_tts" => ModelFamily::GlmTts,
"higgs_audio_tts" => ModelFamily::HiggsAudioTts,
"index_tts2" => ModelFamily::IndexTts2,
"irodori_tts" => ModelFamily::IrodoriTts,
"moss_tts_local" => ModelFamily::MossTtsLocal,
"moss_tts_nano" => ModelFamily::MossTtsNano,
"moss_voicegen" => ModelFamily::MossVoicegen,
"neutts" => ModelFamily::Neutts,
"outetts" => ModelFamily::Outetts,
"pocket_tts" => ModelFamily::PocketTts,
"vietneu_tts" => ModelFamily::VietneuTts,
"fireredtts3" => ModelFamily::Fireredtts3,
"f5_tts" => ModelFamily::F5Tts,
"magpie_tts" => ModelFamily::MagpieTts,
"personaplex" => ModelFamily::Personaplex,
"soprano_tts" => ModelFamily::SopranoTts,
"echo_tts" => ModelFamily::EchoTts,
"voxcpm1" => ModelFamily::Voxcpm1,
"kokoro_tts" => ModelFamily::KokoroTts,
"minimax_h3" => ModelFamily::MinimaxH3,
"sanotts" => ModelFamily::Sanotts,
"sopro_tts" => ModelFamily::SoproTts,
"sopro" => ModelFamily::SoproTts,
"sopro_v2" => ModelFamily::SoproTts,
"sopro_v2_turbo" => ModelFamily::SoproTts,
"mira_tts" => ModelFamily::MiraTts,
"miratts" => ModelFamily::MiraTts,
"mira" => ModelFamily::MiraTts,
"MiraTTS" => ModelFamily::MiraTts,
"cosyvoice3" => ModelFamily::Cosyvoice3,
"breeze_tts" => ModelFamily::BreezeTts,
"sortformer_diar" => ModelFamily::SortformerDiar,
"sortformer_diar_v2" => ModelFamily::SortformerDiarV2,
"seed_vc" => ModelFamily::SeedVc,
"rvc" => ModelFamily::Rvc,
"meanvc2" => ModelFamily::Meanvc2,
"chatterbox" => ModelFamily::Chatterbox,
"chatterbox_turbo" => ModelFamily::ChatterboxTurbo,
"vevo2" => ModelFamily::Vevo2,
"voxcpm2" => ModelFamily::Voxcpm2,
"ace_step" => ModelFamily::AceStep,
"htdemucs" => ModelFamily::Htdemucs,
"mel_band_roformer" => ModelFamily::MelBandRoformer,
"bs_roformer" => ModelFamily::BsRoformer,
"muscriptor" => ModelFamily::Muscriptor,
"omnivoice" => ModelFamily::Omnivoice,
"stable_audio" => ModelFamily::StableAudio,
"minimax_music3" => ModelFamily::MinimaxMusic3,
"supertonic" => ModelFamily::Supertonic,
"voxtral_realtime" => ModelFamily::VoxtralRealtime,
"audiosr" => ModelFamily::Audiosr,
"builtin_audio_utils" => ModelFamily::BuiltinAudioUtils,
"controlfoley" => ModelFamily::ControlFoley,
"midashenglm_gen" => ModelFamily::MidashEnglmGen,
"miocodec" => ModelFamily::Miocodec,
"miotts" => ModelFamily::Miotts,
"vibevoice" => ModelFamily::Vibevoice,
"dramabox" => ModelFamily::Dramabox,
"heartmula" => ModelFamily::Heartmula,
"inflect_v2" => ModelFamily::InflectV2,
"qwen3_forced_aligner" => ModelFamily::Qwen3ForcedAligner,
"mms_forced_aligner" => ModelFamily::MmsForcedAligner,
other => ModelFamily::Custom(other.to_owned()),
}
}
}
impl From<String> for ModelFamily {
fn from(s: String) -> Self {
ModelFamily::from(s.as_str())
}
}
impl AsRef<str> for ModelFamily {
fn as_ref(&self) -> &str {
self.as_str()
}
}
impl std::fmt::Display for ModelFamily {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(self.as_str())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Backend {
Cpu,
Cuda,
Hip,
Vulkan,
Metal,
Best,
}
impl Backend {
pub fn as_str(self) -> &'static str {
match self {
Backend::Cpu => "cpu",
Backend::Cuda => "cuda",
Backend::Hip => "hip",
Backend::Vulkan => "vulkan",
Backend::Metal => "metal",
Backend::Best => "best",
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Device {
pub backend: String,
pub index: i32,
pub name: String,
#[serde(rename = "type")]
pub kind: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct LoaderInfo {
pub family: String,
pub capabilities: Capabilities,
pub instructions_policy: String,
pub api_endpoints: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Capabilities {
pub supported_tasks: Vec<SupportedTask>,
pub languages: Vec<String>,
pub supports_speaker_reference: bool,
pub supports_style_condition: bool,
pub supports_timestamps: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SupportedTask {
pub task: String,
pub modes: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelMetadata {
pub family: String,
pub variant: String,
pub description: String,
pub config_candidates: Vec<String>,
pub weight_candidates: Vec<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct TimeSpan {
pub start_sample: i64,
pub end_sample: i64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SpeechSegment {
pub span: TimeSpan,
pub confidence: f32,
pub text: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SpeakerTurn {
pub speaker_id: String,
pub span: TimeSpan,
pub confidence: f32,
#[serde(default)]
pub text: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct WordTimestamp {
pub span: TimeSpan,
pub word: String,
pub confidence: f32,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VoiceArtifact {
pub kind: String,
pub id: String,
#[serde(default)]
pub payload_base64: String,
#[serde(default)]
pub meta: BTreeMap<String, String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CliOption {
pub name: String,
pub value_name: String,
pub description: String,
pub required: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_value: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub min_value: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_value: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct NamedAsset {
pub id: String,
pub path: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct CliInterface {
#[serde(default)]
pub request_options: Vec<CliOption>,
#[serde(default)]
pub session_options: Vec<CliOption>,
#[serde(default)]
pub load_options: Vec<CliOption>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelInspection {
pub metadata: ModelMetadata,
pub capabilities: Capabilities,
pub cli: CliInterface,
#[serde(default)]
pub discovered_configs: Vec<NamedAsset>,
#[serde(default)]
pub discovered_weights: Vec<NamedAsset>,
#[serde(default)]
pub model_root: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TextOutput {
pub text: String,
pub language: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct AudioBufferInfo {
pub sample_rate: i32,
pub channels: i32,
pub sample_count: usize,
#[serde(default)]
pub samples: Option<Vec<f32>>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct NamedAudioOutput {
pub id: String,
pub audio: AudioBufferInfo,
pub meta: BTreeMap<String, String>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct TaskResult {
pub speech_segments: Vec<SpeechSegment>,
pub speaker_turns: Vec<SpeakerTurn>,
pub text_output: Option<TextOutput>,
pub audio_output: Option<AudioBufferInfo>,
pub named_audio_outputs: Vec<NamedAudioOutput>,
#[serde(default)]
pub word_timestamps: Vec<WordTimestamp>,
#[serde(default)]
pub artifact_output: Option<VoiceArtifact>,
#[serde(default)]
pub output_artifacts: Vec<VoiceArtifact>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct StreamEvent {
pub voice_activity: Vec<VoiceActivityEvent>,
pub partial_text: Option<TextOutput>,
pub audio_output: Option<AudioBufferInfo>,
#[serde(default)]
pub named_audio_outputs: Vec<NamedAudioOutput>,
#[serde(default)]
pub speaker_turns: Vec<SpeakerTurn>,
#[serde(default)]
pub word_timestamps: Vec<WordTimestamp>,
#[serde(default)]
pub output_artifacts: Vec<VoiceArtifact>,
pub is_final: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VoiceActivityEvent {
pub kind: String,
pub sample: i64,
pub probability: f32,
pub segment: Option<SpeechSegment>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct StreamingPolicy {
pub input: String,
pub output: String,
pub preferred_audio_chunk_samples: usize,
pub preferred_audio_chunk_seconds: f64,
}
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used, clippy::expect_used)]
use super::*;
#[test]
fn model_family_roundtrip() {
let mut variants: Vec<ModelFamily> = Vec::new();
macro_rules! push {
($($v:ident),* $(,)?) => {
$(variants.push(ModelFamily::$v);)*
};
}
push![
SileroVad,
MarblenetVad,
Qwen3Asr,
CitrinetAsr,
SenseAsr,
FunAsrNano,
HiggsAudioStt,
HviskeAsr,
KrokoAsr,
NemotronAsr,
ParakeetTdt,
VibevoiceAsr,
VibevoiceAsrStreaming,
FireredAudio,
Granite5Asr,
Audio8Asr,
Audio8Tts,
Vibeasr,
MoonshineAsr,
NiagaraAsr,
Qwen3Tts,
Confucius4Tts,
DotsTts,
FishAudio,
GlmTts,
HiggsAudioTts,
IndexTts2,
IrodoriTts,
MossTtsLocal,
MossTtsNano,
MossVoicegen,
Neutts,
Outetts,
PocketTts,
VietneuTts,
Fireredtts3,
F5Tts,
MagpieTts,
Personaplex,
SopranoTts,
EchoTts,
Voxcpm1,
KokoroTts,
MinimaxH3,
Sanotts,
SoproTts,
MiraTts,
Cosyvoice3,
BreezeTts,
SortformerDiar,
SortformerDiarV2,
SeedVc,
Rvc,
Meanvc2,
Chatterbox,
ChatterboxTurbo,
Vevo2,
Voxcpm2,
AceStep,
Htdemucs,
MelBandRoformer,
BsRoformer,
Muscriptor,
Omnivoice,
StableAudio,
MinimaxMusic3,
Supertonic,
VoxtralRealtime,
Audiosr,
BuiltinAudioUtils,
ControlFoley,
MidashEnglmGen,
Miocodec,
Miotts,
Vibevoice,
Dramabox,
Heartmula,
InflectV2,
Qwen3ForcedAligner,
MmsForcedAligner,
];
assert!(
variants.len() >= 40,
"ModelFamily 应至少覆盖上游全部 40+ loader 族,当前 {}",
variants.len()
);
for v in variants {
let name = v.as_str();
assert_eq!(
ModelFamily::from(name),
v,
"as_str() 与 From<&str> 不一致:{}",
name
);
assert_eq!(format!("{v}"), name);
assert_eq!(v.as_ref(), name);
}
}
#[test]
fn model_family_custom_fallback() {
let m = ModelFamily::from("brand_new_family");
match &m {
ModelFamily::Custom(s) => assert_eq!(s, "brand_new_family"),
other => panic!("未收录名字应落入 Custom,得到 {other:?}"),
}
assert_eq!(m.as_str(), "brand_new_family");
assert_eq!(
ModelFamily::from("brand_new_family"),
ModelFamily::Custom("brand_new_family".to_owned())
);
assert_eq!(
ModelFamily::from("qwen3_asr".to_owned()),
ModelFamily::Qwen3Asr
);
}
#[test]
fn model_family_from_path() {
assert_eq!(
ModelFamily::from_path("models/qwen3-asr-0.6b-q8_0.gguf"),
Some(ModelFamily::Qwen3Asr)
);
assert_eq!(
ModelFamily::from_path("sortformer-diar-4spk-v1-q8_0.gguf"),
Some(ModelFamily::SortformerDiar)
);
assert_eq!(
ModelFamily::from_path("citrinet-asr-q8_0.gguf"),
Some(ModelFamily::CitrinetAsr)
);
assert_eq!(
ModelFamily::from_path("moss-tts-nano-q8_0.gguf"),
Some(ModelFamily::MossTtsNano)
);
assert_eq!(
ModelFamily::from_path("qwen3-tts-12hz-0.6b-base-q8_0.gguf"),
Some(ModelFamily::Qwen3Tts)
);
assert_eq!(
ModelFamily::from_path("qwen3-forced-aligner-q8_0.gguf"),
Some(ModelFamily::Qwen3ForcedAligner)
);
assert_eq!(
ModelFamily::from_path("mira-tts-q8_0.gguf"),
Some(ModelFamily::MiraTts)
);
assert_eq!(
ModelFamily::from_path("htdemucs-6s-q8_0.gguf"),
Some(ModelFamily::Htdemucs)
);
assert_eq!(
ModelFamily::from_path("vibevoice-asr-streaming-7b-q8_0.gguf"),
Some(ModelFamily::VibevoiceAsrStreaming)
);
assert_eq!(
ModelFamily::from_path("sortformer-diar-v2-q8_0.gguf"),
Some(ModelFamily::SortformerDiarV2)
);
assert_eq!(
ModelFamily::from_path("kokoro-tts-82m-q8_0.gguf"),
Some(ModelFamily::KokoroTts)
);
assert_eq!(
ModelFamily::from_path("moonshine-tiny-q8_0.gguf"),
Some(ModelFamily::MoonshineAsr)
);
assert_eq!(
ModelFamily::from_path("niagara-asr-q8_0.gguf"),
Some(ModelFamily::NiagaraAsr)
);
assert_eq!(
ModelFamily::from_path("deepfilternet2"),
Some(ModelFamily::BuiltinAudioUtils)
);
assert_eq!(
ModelFamily::from_path("Qwen3-ASR.Q8_0.GGUF"),
Some(ModelFamily::Qwen3Asr)
);
assert_eq!(
ModelFamily::from_path("model.gguf#family=citrinet_asr"),
Some(ModelFamily::CitrinetAsr)
);
assert_eq!(
ModelFamily::from_path("model.gguf#family=MiraTTS"),
Some(ModelFamily::MiraTts)
);
assert_eq!(
ModelFamily::from_path("model.gguf#family=qwen3_tts"),
Some(ModelFamily::Qwen3Tts)
);
assert_eq!(ModelFamily::from_path("my-weights.gguf"), None);
assert_eq!(ModelFamily::from_path(""), None);
}
#[test]
fn task_kind_as_str() {
let cases = [
(TaskKind::Vad, "vad"),
(TaskKind::Asr, "asr"),
(TaskKind::Diar, "diar"),
(TaskKind::SourceSeparation, "sep"),
(TaskKind::AudioGeneration, "gen"),
(TaskKind::Tts, "tts"),
(TaskKind::VoiceCloning, "clon"),
(TaskKind::VoiceConversion, "vc"),
(TaskKind::SpeechToSpeech, "s2s"),
(TaskKind::Alignment, "align"),
(TaskKind::VoiceDesign, "vdes"),
(TaskKind::SpeakerRecognition, "spk"),
(TaskKind::Svc, "svc"),
(TaskKind::Midi, "midi"),
];
for (k, want) in cases {
assert_eq!(k.as_str(), want);
}
}
#[test]
fn task_kind_serde_roundtrip() {
let all = [
TaskKind::Vad,
TaskKind::Asr,
TaskKind::Diar,
TaskKind::SourceSeparation,
TaskKind::AudioGeneration,
TaskKind::Tts,
TaskKind::VoiceCloning,
TaskKind::VoiceConversion,
TaskKind::SpeechToSpeech,
TaskKind::Alignment,
TaskKind::VoiceDesign,
TaskKind::SpeakerRecognition,
TaskKind::Svc,
TaskKind::Midi,
];
for k in all {
let s = serde_json::to_string(&k).unwrap();
assert_eq!(s, format!("\"{}\"", k.as_str()));
assert_eq!(serde_json::from_str::<TaskKind>(&s).unwrap(), k);
}
assert!(serde_json::from_str::<TaskKind>("\"sourceseparation\"").is_err());
}
#[test]
fn run_mode_as_str() {
assert_eq!(RunMode::Offline.as_str(), "offline");
assert_eq!(RunMode::Streaming.as_str(), "streaming");
}
#[test]
fn backend_as_str() {
let cases = [
(Backend::Cpu, "cpu"),
(Backend::Cuda, "cuda"),
(Backend::Hip, "hip"),
(Backend::Vulkan, "vulkan"),
(Backend::Metal, "metal"),
(Backend::Best, "best"),
];
for (b, want) in cases {
assert_eq!(b.as_str(), want);
}
}
#[test]
fn structured_types_deserialize() {
let result: TaskResult = serde_json::from_str(
r#"{
"speech_segments": [
{"span": {"start_sample": 0, "end_sample": 1600}, "confidence": 0.95, "text": ""}
],
"text_output": {"text": "hi", "language": "en"},
"audio_output": {"sample_rate": 24000, "channels": 1, "sample_count": 0, "samples": []},
"named_audio_outputs": [],
"speaker_turns": [
{"speaker_id": "speaker_0", "span": {"start_sample": 0, "end_sample": 1600}, "confidence": 0.9, "text": "hello"}
],
"word_timestamps": [
{"span": {"start_sample": 0, "end_sample": 320}, "word": "hello", "confidence": 0.99}
],
"artifact_output": {"kind": "speaker_embedding", "id": "spk1", "payload_base64": "", "meta": {}},
"output_artifacts": [
{"kind": "custom", "id": "a", "payload_base64": "AQID", "meta": {"k": "v"}}
]
}"#,
)
.unwrap();
assert_eq!(result.speech_segments.len(), 1);
assert_eq!(result.speech_segments[0].span.start_sample, 0);
assert_eq!(result.speech_segments[0].confidence, 0.95);
assert_eq!(result.text_output.as_ref().unwrap().text, "hi");
assert_eq!(result.audio_output.unwrap().sample_rate, 24000);
assert_eq!(result.speaker_turns[0].text, "hello");
assert_eq!(result.word_timestamps[0].word, "hello");
assert_eq!(
result.artifact_output.as_ref().unwrap().kind,
"speaker_embedding"
);
assert_eq!(result.output_artifacts[0].payload_base64, "AQID");
let ev: StreamEvent = serde_json::from_str(
r#"{
"voice_activity": [
{"kind": "speech_start", "sample": 100, "probability": 0.8, "segment": null}
],
"partial_text": null,
"audio_output": null,
"named_audio_outputs": [{"id": "chunk_0", "audio": {"sample_rate": 48000, "channels": 2, "sample_count": 320, "samples": []}, "meta": {}}],
"speaker_turns": [],
"word_timestamps": [],
"output_artifacts": [],
"is_final": true
}"#,
)
.unwrap();
assert_eq!(ev.voice_activity[0].kind, "speech_start");
assert_eq!(ev.named_audio_outputs[0].id, "chunk_0");
assert!(ev.is_final);
}
#[test]
fn model_inspection_deserialize() {
let info: ModelInspection = serde_json::from_str(
r#"{
"metadata": {
"family": "qwen3_asr",
"variant": "q8_0",
"description": "Qwen3 ASR",
"config_candidates": ["config.json"],
"weight_candidates": ["qwen3-asr-q8_0.gguf"]
},
"capabilities": {
"supported_tasks": [{"task": "asr", "modes": ["offline", "streaming"]}],
"languages": ["zh", "en"],
"supports_speaker_reference": false,
"supports_style_condition": false,
"supports_timestamps": true
},
"discovered_weights": [
{"id": "model", "path": "/models/qwen3-asr-q8_0.gguf"}
],
"cli": {
"request_options": [
{"name": "language", "value_name": "CODE", "description": "语言", "required": false, "default_value": "auto"}
],
"session_options": [],
"load_options": []
}
}"#,
)
.unwrap();
assert_eq!(info.metadata.family, "qwen3_asr");
assert!(
info.capabilities
.supported_tasks
.iter()
.any(|t| t.task == "asr" && t.modes.contains(&"streaming".to_string()))
);
assert_eq!(info.discovered_weights[0].id, "model");
assert_eq!(info.cli.request_options[0].name, "language");
assert_eq!(
info.cli.request_options[0].default_value.as_deref(),
Some("auto")
);
}
}