use serde::{Deserialize, Serialize};
use crate::models::context_fit::{self, ContextFit};
use crate::registry::ConfigConstructable;
use crate::utils::Searchable;
use crate::utils::hardware::HardwareProfile;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize, schemars::JsonSchema)]
pub enum ModelFunction {
Chat,
ToolCalling,
Thinking,
ImageUnderstanding,
Guardian,
Embeddings,
Transcription,
Translation,
SpeakerAttribution,
KeywordBiasing,
}
impl std::fmt::Display for ModelFunction {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ModelFunction::Chat => write!(f, "Chat"),
ModelFunction::ToolCalling => write!(f, "ToolCalling"),
ModelFunction::Thinking => write!(f, "Thinking"),
ModelFunction::ImageUnderstanding => write!(f, "Image Understanding"),
ModelFunction::Guardian => write!(f, "Guardian"),
ModelFunction::Embeddings => write!(f, "Embeddings"),
ModelFunction::Transcription => write!(f, "Transcription"),
ModelFunction::Translation => write!(f, "Translation"),
ModelFunction::SpeakerAttribution => write!(f, "Speaker Attribution"),
ModelFunction::KeywordBiasing => write!(f, "Keyword Biasing"),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum LayerKind {
FullAttention,
SlidingAttention { window: u64 },
Recurrent(MambaShape),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MambaShape {
pub d_conv: u64,
pub d_state: u64,
pub d_inner: u64,
pub n_groups: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LayerTypeCount {
pub kind: LayerKind,
pub count: u64,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelArchitecture {
pub num_hidden_layers: u64,
pub hidden_size: u64,
pub num_attention_heads: u64,
pub num_key_value_heads: u64,
pub head_dim: u64,
pub layer_types: Vec<LayerTypeCount>,
}
pub trait Model: ConfigConstructable {
fn family(&self) -> &str;
fn version(&self) -> &str;
fn size(&self) -> u64;
fn context_length(&self) -> u64;
fn model_type(&self) -> &ModelType;
fn huggingface_repo(&self) -> &str;
fn native_dtype(&self) -> &str;
fn architecture(&self) -> &ModelArchitecture;
fn variants(&self) -> &[ModelVariant];
fn description(&self) -> Option<&str>;
fn tags(&self) -> &[String];
fn supported_functions(&self) -> &[ModelFunction];
fn context_fit(&self, variant: &ModelVariant, hardware: &HardwareProfile) -> ContextFit {
context_fit::estimate(
self.context_length(),
self.architecture(),
self.native_dtype(),
variant,
hardware,
)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelMetadata {
pub family: String,
pub version: String,
pub size: u64,
pub context_length: u64,
pub model_type: ModelType,
pub huggingface_repo: String,
pub native_dtype: String,
pub architecture: ModelArchitecture,
pub variants: Vec<ModelVariant>,
pub description: Option<String>,
pub tags: Vec<String>,
pub supported_functions: Vec<ModelFunction>,
}
impl ModelMetadata {
pub fn format_size(&self) -> String {
if self.size >= 1_000_000_000 {
format!("{}B", self.size / 1_000_000_000)
} else {
format!("{}M", self.size / 1_000_000)
}
}
}
impl std::fmt::Display for ModelMetadata {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"{}) - {} params, {} context, Type: {}",
self.family,
self.format_size(),
self.context_length,
self.model_type
)
}
}
impl Searchable for ModelMetadata {
fn search_fields(&self) -> Vec<&str> {
let mut fields: Vec<&str> = vec![self.family.as_str()];
if let Some(desc) = &self.description {
fields.push(desc.as_str());
}
fields.extend(self.tags.iter().map(String::as_str));
fields
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum ModelType {
Text,
Vision,
Speech,
Embedding,
}
impl std::fmt::Display for ModelType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ModelType::Text => write!(f, "Text"),
ModelType::Vision => write!(f, "Vision"),
ModelType::Speech => write!(f, "Speech"),
ModelType::Embedding => write!(f, "Embedding"),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelVariant {
pub format: String,
pub precision: String,
pub size_gb: f64,
pub url: String,
}
use crate::define_factory;
define_factory!(Model, ModelMetadata, ModelFactory);
#[cfg(test)]
mod format_size_tests {
use super::*;
fn test_architecture() -> ModelArchitecture {
ModelArchitecture {
num_hidden_layers: 32,
hidden_size: 4096,
num_attention_heads: 32,
num_key_value_heads: 8,
head_dim: 128,
layer_types: vec![LayerTypeCount {
kind: LayerKind::FullAttention,
count: 32,
}],
}
}
fn metadata_with_size(size: u64) -> ModelMetadata {
ModelMetadata {
family: "Test".to_string(),
version: "1.0".to_string(),
size,
context_length: 4096,
model_type: ModelType::Text,
huggingface_repo: "test/test".to_string(),
native_dtype: "bfloat16".to_string(),
architecture: test_architecture(),
variants: vec![],
description: None,
tags: vec![],
supported_functions: vec![],
}
}
#[test]
fn format_size_billions() {
assert_eq!(metadata_with_size(8_000_000_000).format_size(), "8B");
}
#[test]
fn format_size_millions() {
assert_eq!(metadata_with_size(258_000_000).format_size(), "258M");
}
#[test]
fn format_size_boundary_is_one_billion() {
assert_eq!(metadata_with_size(1_000_000_000).format_size(), "1B");
assert_eq!(metadata_with_size(999_999_999).format_size(), "999M");
}
#[test]
fn format_size_30m_model() {
assert_eq!(metadata_with_size(30_295_296).format_size(), "30M");
}
}
#[cfg(test)]
mod searchable_tests {
use super::*;
fn metadata(family: &str, description: Option<&str>, tags: Vec<&str>) -> ModelMetadata {
ModelMetadata {
family: family.to_string(),
version: "1.0".to_string(),
size: 8_000_000_000,
context_length: 4096,
model_type: ModelType::Text,
huggingface_repo: "ibm-granite/test".to_string(),
native_dtype: "bfloat16".to_string(),
architecture: ModelArchitecture {
num_hidden_layers: 32,
hidden_size: 4096,
num_attention_heads: 32,
num_key_value_heads: 8,
head_dim: 128,
layer_types: vec![LayerTypeCount {
kind: LayerKind::FullAttention,
count: 32,
}],
},
variants: vec![],
description: description.map(String::from),
tags: tags.into_iter().map(String::from).collect(),
supported_functions: vec![],
}
}
#[test]
fn searchable_fields_includes_family() {
let m = metadata("Granite 3.1", None, vec![]);
assert!(m.search_fields().contains(&"Granite 3.1"));
}
#[test]
fn searchable_fields_includes_description_when_present() {
let m = metadata("Granite 3.1", Some("A text model"), vec![]);
assert!(m.search_fields().contains(&"A text model"));
}
#[test]
fn searchable_fields_omits_description_when_absent() {
let m = metadata("Granite 3.1", None, vec![]);
assert_eq!(m.search_fields().len(), 1);
}
#[test]
fn searchable_fields_includes_tags() {
let m = metadata("Granite 3.1", None, vec!["instruct", "chat"]);
let fields = m.search_fields();
assert!(fields.contains(&"instruct"));
assert!(fields.contains(&"chat"));
}
}