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("qwen2", |source, device| {
47 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
48 });
49 r.register("qwen3", |source, device| {
52 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
53 });
54 r.register("mistral", |source, device| {
58 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
59 });
60 r.register("phi3", |source, device| {
65 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
66 });
67 r.register("idefics3", |source, device| {
69 Ok(Box::new(crate::smolvlm::SmolVlmModel::<B>::load(
70 source, device,
71 )?))
72 });
73 r.register("gemma3_text", |source, device| {
79 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
80 });
81 r.register("gemma3", |source, device| {
82 Ok(Box::new(crate::llama::LlamaModel::<B>::load(source, device)?))
83 });
84 r
85 }
86
87 pub fn register(&mut self, architecture: &str, loader: Loader<B>) {
89 self.loaders.insert(architecture.to_string(), loader);
90 }
91
92 pub fn supports(&self, architecture: &str) -> bool {
94 self.loaders.contains_key(architecture)
95 }
96
97 pub fn architectures(&self) -> Vec<&str> {
99 let mut v: Vec<&str> = self.loaders.keys().map(String::as_str).collect();
100 v.sort();
101 v
102 }
103
104 pub fn load(
106 &self,
107 source: &dyn ModelSource,
108 device: &Device<B>,
109 ) -> Result<Box<dyn GenerativeModel<B>>> {
110 let arch = &source.metadata().architecture;
111 let loader = self
112 .loaders
113 .get(arch)
114 .ok_or_else(|| ModelError::UnsupportedArchitecture(arch.clone()))?;
115 loader(source, device)
116 }
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122
123 #[test]
124 fn registry_has_llama_aliases() {
125 let r = ModelRegistry::<burn::backend::NdArray<f32>>::new();
126 assert!(r.supports("llama"));
127 assert!(r.supports("smollm2"));
128 assert!(r.supports("qwen2"));
129 assert!(r.supports("mistral"));
130 assert!(r.supports("qwen3"));
131 }
132}