use std::collections::HashMap;
use burn::tensor::{Device, backend::Backend};
use combs_formats::ModelSource;
use crate::traits::GenerativeModel;
use crate::{ModelError, Result};
pub type Loader<B> =
fn(&dyn ModelSource, &Device<B>) -> Result<Box<dyn GenerativeModel<B>>>;
pub struct ModelRegistry<B: Backend> {
loaders: HashMap<String, Loader<B>>,
}
impl<B: Backend> Default for ModelRegistry<B> {
fn default() -> Self {
Self::new()
}
}
impl<B: Backend> ModelRegistry<B> {
pub fn new() -> Self {
let mut r = ModelRegistry {
loaders: HashMap::new(),
};
r.register("llama", |source, device| {
Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
});
r.register("smollm2", |source, device| {
Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
});
r.register("idefics3", |source, device| {
Ok(Box::new(crate::smolvlm::SmolVlmModel::<B>::load(
source, device,
)?))
});
r
}
pub fn register(&mut self, architecture: &str, loader: Loader<B>) {
self.loaders.insert(architecture.to_string(), loader);
}
pub fn supports(&self, architecture: &str) -> bool {
self.loaders.contains_key(architecture)
}
pub fn architectures(&self) -> Vec<&str> {
let mut v: Vec<&str> = self.loaders.keys().map(String::as_str).collect();
v.sort();
v
}
pub fn load(
&self,
source: &dyn ModelSource,
device: &Device<B>,
) -> Result<Box<dyn GenerativeModel<B>>> {
let arch = &source.metadata().architecture;
let loader = self
.loaders
.get(arch)
.ok_or_else(|| ModelError::UnsupportedArchitecture(arch.clone()))?;
loader(source, device)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn registry_has_llama_aliases() {
let r = ModelRegistry::<burn::backend::NdArray<f32>>::new();
assert!(r.supports("llama"));
assert!(r.supports("smollm2"));
assert!(!r.supports("qwen3"));
}
}