use std::collections::BTreeSet;
use std::path::Path;
use crate::records::{Capability, ExecutionMode, JsonValue, Modality};
use crate::resolution::pipelines::{PipelineFamilyRegistry, diffusers_pipeline_class};
#[derive(Debug, Clone, PartialEq)]
pub struct Hint {
pub modality: Option<Modality>,
pub capabilities: Vec<Capability>,
pub execution: ExecutionMode,
pub context_length: Option<i64>,
pub quantization: Option<String>,
}
impl Hint {
fn new(modality: Modality, capabilities: Vec<Capability>, execution: ExecutionMode) -> Self {
Self {
modality: Some(modality),
capabilities,
execution,
context_length: None,
quantization: None,
}
}
pub fn unknown(execution: ExecutionMode) -> Self {
Self {
modality: None,
capabilities: Vec::new(),
execution,
context_length: None,
quantization: None,
}
}
}
pub fn speech_hint() -> Hint {
Hint::new(
Modality::speech(),
vec![Capability::speak()],
ExecutionMode::Stream,
)
}
pub fn audio_hint() -> Hint {
Hint::new(
Modality::audio(),
vec![Capability::transcribe()],
ExecutionMode::Stream,
)
}
pub fn text_hint() -> Hint {
Hint::new(
Modality::text(),
vec![Capability::chat(), Capability::complete()],
ExecutionMode::Stream,
)
}
pub fn embedding_hint() -> Hint {
Hint::new(
Modality::embedding(),
vec![Capability::embed()],
ExecutionMode::Stream,
)
}
pub fn vision_chat_hint() -> Hint {
Hint::new(
Modality::text(),
vec![
Capability::chat(),
Capability::complete(),
Capability::see(),
],
ExecutionMode::Stream,
)
}
pub fn gguf_hint() -> Hint {
text_hint()
}
pub fn whisper_bin_hint() -> Hint {
audio_hint()
}
pub fn from_model_index(path: &Path) -> Hint {
if let Some(class) = diffusers_pipeline_class(path)
&& let Some(family) = PipelineFamilyRegistry::shared().family(&class)
{
return Hint {
modality: Some(family.modality.clone()),
capabilities: family.capabilities.clone(),
execution: ExecutionMode::Job,
context_length: None,
quantization: None,
};
}
Hint::unknown(ExecutionMode::Job)
}
const VISION_LANGUAGE_ARCHITECTURES: [&str; 5] =
["Llava", "Qwen2VL", "Idefics", "PaliGemma", "Mllama"];
pub fn from_config_json(path: &Path) -> Option<Hint> {
let bytes = std::fs::read(path).ok()?;
let json = serde_json::from_slice::<JsonValue>(&bytes).ok()?;
from_config(&json)
}
pub fn from_config(json: &JsonValue) -> Option<Hint> {
let object = json.as_object()?;
let architectures: Vec<&str> = object
.get("architectures")
.and_then(JsonValue::as_array)
.map(|items| items.iter().filter_map(JsonValue::as_str).collect())
.unwrap_or_default();
let context_length = ["max_position_embeddings", "n_positions", "max_seq_len"]
.into_iter()
.find_map(|key| object.get(key).and_then(JsonValue::as_i64))
.filter(|value| *value > 0);
let quantization = object
.get("quantization")
.and_then(JsonValue::as_object)
.and_then(|block| block.get("bits"))
.and_then(JsonValue::as_i64)
.filter(|bits| *bits > 0)
.map(|bits| format!("{bits}bit"));
if object.contains_key("vision_config")
&& architectures.iter().any(|architecture| {
architecture.ends_with("ForConditionalGeneration")
|| VISION_LANGUAGE_ARCHITECTURES
.iter()
.any(|marker| architecture.contains(marker))
})
{
return Some(with_facts(
vision_chat_hint(),
context_length,
quantization.clone(),
));
}
for architecture in &architectures {
if let Some(hint) = architecture_hint(architecture) {
return Some(with_facts(hint, context_length, quantization.clone()));
}
}
let keys: BTreeSet<&str> = object.keys().map(String::as_str).collect();
config_key_hint(&keys).map(|hint| with_facts(hint, context_length, quantization))
}
fn with_facts(mut hint: Hint, context_length: Option<i64>, quantization: Option<String>) -> Hint {
hint.context_length = context_length;
hint.quantization = quantization;
hint
}
fn architecture_hint(architecture: &str) -> Option<Hint> {
const SPEECH: [&str; 6] = ["Kokoro", "StyleTTS", "Bark", "ParlerTTS", "Vits", "Xtts"];
const EMBEDDING: [&str; 5] = [
"BertModel",
"NomicBertModel",
"ModernBertModel",
"XLMRobertaModel",
"MPNetModel",
];
if SPEECH.iter().any(|marker| architecture.contains(marker)) {
return Some(speech_hint());
}
if architecture.contains("Whisper") {
return Some(audio_hint());
}
if EMBEDDING.iter().any(|marker| architecture.contains(marker)) {
return Some(embedding_hint());
}
if architecture.contains("LMHead") || architecture.ends_with("ForCausalLM") {
return Some(text_hint());
}
None
}
fn config_key_hint(keys: &BTreeSet<&str>) -> Option<Hint> {
if keys.contains("istftnet") || keys.contains("plbert") {
return Some(speech_hint());
}
if keys.contains("style_dim") && keys.contains("n_mels") {
return Some(speech_hint());
}
None
}