kernel/discovery/
modality_hints.rs1use std::collections::BTreeSet;
8use std::path::Path;
9
10use crate::records::{Capability, ExecutionMode, JsonValue, Modality};
11use crate::resolution::pipelines::{PipelineFamilyRegistry, diffusers_pipeline_class};
12
13#[derive(Debug, Clone, PartialEq)]
15pub struct Hint {
16 pub modality: Option<Modality>,
18 pub capabilities: Vec<Capability>,
20 pub execution: ExecutionMode,
22 pub context_length: Option<i64>,
24}
25
26impl Hint {
27 fn new(modality: Modality, capabilities: Vec<Capability>, execution: ExecutionMode) -> Self {
28 Self {
29 modality: Some(modality),
30 capabilities,
31 execution,
32 context_length: None,
33 }
34 }
35
36 pub fn unknown(execution: ExecutionMode) -> Self {
39 Self {
40 modality: None,
41 capabilities: Vec::new(),
42 execution,
43 context_length: None,
44 }
45 }
46}
47
48pub fn speech_hint() -> Hint {
50 Hint::new(
51 Modality::speech(),
52 vec![Capability::speak()],
53 ExecutionMode::Stream,
54 )
55}
56
57pub fn audio_hint() -> Hint {
59 Hint::new(
60 Modality::audio(),
61 vec![Capability::transcribe()],
62 ExecutionMode::Stream,
63 )
64}
65
66pub fn text_hint() -> Hint {
68 Hint::new(
69 Modality::text(),
70 vec![Capability::chat(), Capability::complete()],
71 ExecutionMode::Stream,
72 )
73}
74
75pub fn embedding_hint() -> Hint {
77 Hint::new(
78 Modality::embedding(),
79 vec![Capability::embed()],
80 ExecutionMode::Stream,
81 )
82}
83
84pub fn vision_chat_hint() -> Hint {
86 Hint::new(
87 Modality::text(),
88 vec![
89 Capability::chat(),
90 Capability::complete(),
91 Capability::see(),
92 ],
93 ExecutionMode::Stream,
94 )
95}
96
97pub fn gguf_hint() -> Hint {
99 text_hint()
100}
101
102pub fn whisper_bin_hint() -> Hint {
104 audio_hint()
105}
106
107pub fn from_model_index(path: &Path) -> Hint {
111 if let Some(class) = diffusers_pipeline_class(path)
112 && let Some(family) = PipelineFamilyRegistry::shared().family(&class)
113 {
114 return Hint {
115 modality: Some(family.modality.clone()),
116 capabilities: family.capabilities.clone(),
117 execution: ExecutionMode::Job,
118 context_length: None,
119 };
120 }
121 Hint::unknown(ExecutionMode::Job)
122}
123
124const VISION_LANGUAGE_ARCHITECTURES: [&str; 5] =
127 ["Llava", "Qwen2VL", "Idefics", "PaliGemma", "Mllama"];
128
129pub fn from_config_json(path: &Path) -> Option<Hint> {
132 let bytes = std::fs::read(path).ok()?;
133 let json = serde_json::from_slice::<JsonValue>(&bytes).ok()?;
134 from_config(&json)
135}
136
137pub fn from_config(json: &JsonValue) -> Option<Hint> {
142 let object = json.as_object()?;
143 let architectures: Vec<&str> = object
144 .get("architectures")
145 .and_then(JsonValue::as_array)
146 .map(|items| items.iter().filter_map(JsonValue::as_str).collect())
147 .unwrap_or_default();
148
149 let context_length = ["max_position_embeddings", "n_positions", "max_seq_len"]
152 .into_iter()
153 .find_map(|key| object.get(key).and_then(JsonValue::as_i64))
154 .filter(|value| *value > 0);
155
156 if object.contains_key("vision_config")
157 && architectures.iter().any(|architecture| {
158 architecture.ends_with("ForConditionalGeneration")
159 || VISION_LANGUAGE_ARCHITECTURES
160 .iter()
161 .any(|marker| architecture.contains(marker))
162 })
163 {
164 return Some(with_context(vision_chat_hint(), context_length));
165 }
166
167 for architecture in &architectures {
168 if let Some(hint) = architecture_hint(architecture) {
169 return Some(with_context(hint, context_length));
170 }
171 }
172
173 let keys: BTreeSet<&str> = object.keys().map(String::as_str).collect();
174 config_key_hint(&keys).map(|hint| with_context(hint, context_length))
175}
176
177fn with_context(mut hint: Hint, context_length: Option<i64>) -> Hint {
178 hint.context_length = context_length;
179 hint
180}
181
182fn architecture_hint(architecture: &str) -> Option<Hint> {
185 const SPEECH: [&str; 6] = ["Kokoro", "StyleTTS", "Bark", "ParlerTTS", "Vits", "Xtts"];
186 const EMBEDDING: [&str; 5] = [
187 "BertModel",
188 "NomicBertModel",
189 "ModernBertModel",
190 "XLMRobertaModel",
191 "MPNetModel",
192 ];
193 if SPEECH.iter().any(|marker| architecture.contains(marker)) {
194 return Some(speech_hint());
195 }
196 if architecture.contains("Whisper") {
197 return Some(audio_hint());
198 }
199 if EMBEDDING.iter().any(|marker| architecture.contains(marker)) {
200 return Some(embedding_hint());
201 }
202 if architecture.contains("LMHead") || architecture.ends_with("ForCausalLM") {
203 return Some(text_hint());
204 }
205 None
206}
207
208fn config_key_hint(keys: &BTreeSet<&str>) -> Option<Hint> {
211 if keys.contains("istftnet") || keys.contains("plbert") {
212 return Some(speech_hint());
213 }
214 if keys.contains("style_dim") && keys.contains("n_mels") {
215 return Some(speech_hint());
216 }
217 None
218}