use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum TaskKind {
Vad,
Asr,
Diar,
SourceSeparation,
AudioGeneration,
Tts,
}
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",
}
}
}
#[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,
Qwen3Tts,
Confucius4Tts,
DotsTts,
FishAudio,
GlmTts,
HiggsAudioTts,
IndexTts2,
IrodoriTts,
MossTtsLocal,
MossTtsNano,
Neutts,
Outetts,
PocketTts,
VietneuTts,
MinimaxH3,
SortformerDiar,
SeedVc,
Rvc,
Chatterbox,
Vevo2,
Voxcpm2,
AceStep,
Htdemucs,
MelBandRoformer,
BsRoformer,
Muscriptor,
Omnivoice,
StableAudio,
MinimaxMusic3,
Supertonic,
VoxtralRealtime,
Miocodec,
Miotts,
Vibevoice,
Dramabox,
Heartmula,
InflectV2,
Qwen3ForcedAligner,
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", 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", ModelFamily::VibevoiceAsr),
("vibevoice", ModelFamily::Vibevoice),
("qwen3_tts", ModelFamily::Qwen3Tts),
("qwen3-tts", ModelFamily::Qwen3Tts),
("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", ModelFamily::MossTtsNano),
("neutts", ModelFamily::Neutts),
("outetts", ModelFamily::Outetts),
("pocket_tts", ModelFamily::PocketTts),
("pocket-tts", ModelFamily::PocketTts),
("vietneu", ModelFamily::VietneuTts),
("minimax_h3", ModelFamily::MinimaxH3),
("minimax-h3", ModelFamily::MinimaxH3),
("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),
("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),
("miocodec", ModelFamily::Miocodec),
("miotts", ModelFamily::Miotts),
("dramabox", ModelFamily::Dramabox),
("heartmula", ModelFamily::Heartmula),
("inflect", ModelFamily::InflectV2),
("qwen3_forced_aligner", ModelFamily::Qwen3ForcedAligner),
("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::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::Neutts => "neutts",
ModelFamily::Outetts => "outetts",
ModelFamily::PocketTts => "pocket_tts",
ModelFamily::VietneuTts => "vietneu_tts",
ModelFamily::MinimaxH3 => "minimax_h3",
ModelFamily::SortformerDiar => "sortformer_diar",
ModelFamily::SeedVc => "seed_vc",
ModelFamily::Rvc => "rvc",
ModelFamily::Chatterbox => "chatterbox",
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::Miocodec => "miocodec",
ModelFamily::Miotts => "miotts",
ModelFamily::Vibevoice => "vibevoice",
ModelFamily::Dramabox => "dramabox",
ModelFamily::Heartmula => "heartmula",
ModelFamily::InflectV2 => "inflect_v2",
ModelFamily::Qwen3ForcedAligner => "qwen3_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,
"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,
"neutts" => ModelFamily::Neutts,
"outetts" => ModelFamily::Outetts,
"pocket_tts" => ModelFamily::PocketTts,
"vietneu_tts" => ModelFamily::VietneuTts,
"minimax_h3" => ModelFamily::MinimaxH3,
"sortformer_diar" => ModelFamily::SortformerDiar,
"seed_vc" => ModelFamily::SeedVc,
"rvc" => ModelFamily::Rvc,
"chatterbox" => ModelFamily::Chatterbox,
"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,
"miocodec" => ModelFamily::Miocodec,
"miotts" => ModelFamily::Miotts,
"vibevoice" => ModelFamily::Vibevoice,
"dramabox" => ModelFamily::Dramabox,
"heartmula" => ModelFamily::Heartmula,
"inflect_v2" => ModelFamily::InflectV2,
"qwen3_forced_aligner" => ModelFamily::Qwen3ForcedAligner,
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)]
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,
}
#[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>,
}
#[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>,
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,
Qwen3Tts,
Confucius4Tts,
DotsTts,
FishAudio,
GlmTts,
HiggsAudioTts,
IndexTts2,
IrodoriTts,
MossTtsLocal,
MossTtsNano,
Neutts,
Outetts,
PocketTts,
VietneuTts,
MinimaxH3,
SortformerDiar,
SeedVc,
Rvc,
Chatterbox,
Vevo2,
Voxcpm2,
AceStep,
Htdemucs,
MelBandRoformer,
BsRoformer,
Muscriptor,
Omnivoice,
StableAudio,
Supertonic,
VoxtralRealtime,
Miocodec,
Miotts,
Vibevoice,
Dramabox,
Heartmula,
InflectV2,
Qwen3ForcedAligner,
];
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("htdemucs-6s-q8_0.gguf"),
Some(ModelFamily::Htdemucs)
);
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=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"),
];
for (k, want) in cases {
assert_eq!(k.as_str(), want);
}
}
#[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": []
}"#,
)
.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);
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": {}}],
"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);
}
}