1use std::collections::HashMap;
5
6use burn::tensor::{Device, backend::Backend};
7use combs_formats::ModelSource;
8
9use crate::traits::GenerativeModel;
10use crate::{ModelError, Result};
11
12pub type Loader<B> =
14 fn(&dyn ModelSource, &Device<B>) -> Result<Box<dyn GenerativeModel<B>>>;
15
16pub struct ModelRegistry<B: Backend> {
19 loaders: HashMap<String, Loader<B>>,
20}
21
22impl<B: Backend> Default for ModelRegistry<B> {
23 fn default() -> Self {
24 Self::new()
25 }
26}
27
28impl<B: Backend> ModelRegistry<B> {
29 pub fn new() -> Self {
31 let mut r = ModelRegistry {
32 loaders: HashMap::new(),
33 };
34 r.register("llama", |source, device| {
37 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
38 });
39 r.register("smollm2", |source, device| {
40 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
41 });
42 r.register("idefics3", |source, device| {
44 Ok(Box::new(crate::smolvlm::SmolVlmModel::<B>::load(
45 source, device,
46 )?))
47 });
48 r
49 }
50
51 pub fn register(&mut self, architecture: &str, loader: Loader<B>) {
53 self.loaders.insert(architecture.to_string(), loader);
54 }
55
56 pub fn supports(&self, architecture: &str) -> bool {
58 self.loaders.contains_key(architecture)
59 }
60
61 pub fn architectures(&self) -> Vec<&str> {
63 let mut v: Vec<&str> = self.loaders.keys().map(String::as_str).collect();
64 v.sort();
65 v
66 }
67
68 pub fn load(
70 &self,
71 source: &dyn ModelSource,
72 device: &Device<B>,
73 ) -> Result<Box<dyn GenerativeModel<B>>> {
74 let arch = &source.metadata().architecture;
75 let loader = self
76 .loaders
77 .get(arch)
78 .ok_or_else(|| ModelError::UnsupportedArchitecture(arch.clone()))?;
79 loader(source, device)
80 }
81}
82
83#[cfg(test)]
84mod tests {
85 use super::*;
86
87 #[test]
88 fn registry_has_llama_aliases() {
89 let r = ModelRegistry::<burn::backend::NdArray<f32>>::new();
90 assert!(r.supports("llama"));
91 assert!(r.supports("smollm2"));
92 assert!(!r.supports("qwen3"));
93 }
94}