Skip to main content

combs_models/
registry.rs

1//! Architecture registry: maps `metadata.architecture` to a model loader.
2//! New architectures are additive — one module + one `register` call.
3
4use std::collections::HashMap;
5
6use burn::tensor::{Device, backend::Backend};
7use combs_formats::ModelSource;
8
9use crate::traits::GenerativeModel;
10use crate::{ModelError, Result};
11
12/// A constructor for a boxed model of some architecture.
13pub type Loader<B> =
14    fn(&dyn ModelSource, &Device<B>) -> Result<Box<dyn GenerativeModel<B>>>;
15
16/// Maps architecture identifiers (`config.json::model_type`, plus known
17/// aliases) to loaders. Mirrors MLC's `model.py::MODELS` table.
18pub 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    /// Creates a registry with the built-in architectures registered.
30    pub fn new() -> Self {
31        let mut r = ModelRegistry {
32            loaders: HashMap::new(),
33        };
34        // SmolLM2 reports model_type "llama" (older releases: "smollm2");
35        // both are Llama-structured.
36        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        // SmolVLM reports model_type "idefics3" (SigLIP + pixel-shuffle + SmolLM2).
43        r.register("idefics3", |source, device| {
44            Ok(Box::new(crate::smolvlm::SmolVlmModel::<B>::load(
45                source, device,
46            )?))
47        });
48        r
49    }
50
51    /// Registers (or replaces) the loader for an architecture id.
52    pub fn register(&mut self, architecture: &str, loader: Loader<B>) {
53        self.loaders.insert(architecture.to_string(), loader);
54    }
55
56    /// Whether an architecture id has a loader.
57    pub fn supports(&self, architecture: &str) -> bool {
58        self.loaders.contains_key(architecture)
59    }
60
61    /// Registered architecture ids.
62    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    /// Loads the model described by `source`'s metadata.
69    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}