Skip to main content

combs_models/
archspec.rs

1//! Per-architecture resolution: turns parsed [`ModelMetadata`] into the
2//! concrete knobs the universal decoder consumes — activation, norm
3//! flavor, QK-norm, sandwich norms, embedding scale, logit softcaps, RoPE
4//! configuration, and the per-layer attention layout.
5//!
6//! `combs-formats` stays parse-only; everything HF configs can't express
7//! (gemma's sandwich norms, its `(1+w)` RMSNorm flavor, per-family layout
8//! semantics) is encoded here, keyed on the architecture string. Adding an
9//! architecture = adding a resolver arm, not decoder code.
10
11use combs_formats::{Activation, ModelMetadata, RopeScaling};
12
13/// RMSNorm flavor: plain (`x̂·w`) or gemma's zero-centered (`x̂·(1+w)`).
14#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
15pub enum NormFlavor {
16    #[default]
17    RmsNorm,
18    GemmaRmsNorm,
19}
20
21/// Attention kind of one layer.
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23pub enum LayerKind {
24    Global,
25    Sliding(usize),
26}
27
28/// Resolved architecture description consumed by the decoder.
29#[derive(Debug, Clone)]
30pub struct ArchSpec {
31    pub activation: Activation,
32    pub norm_flavor: NormFlavor,
33    /// Per-head RMSNorm on Q and K after projection (gemma-3, qwen-3).
34    pub qk_norm: bool,
35    /// Extra post-attention / post-MLP norms around the residual adds
36    /// (gemma-3 "sandwich" norms).
37    pub sandwich_norms: bool,
38    /// Multiply embeddings by `sqrt(hidden_size)` (gemma; computed in f32).
39    pub embed_scale_sqrt_hidden: bool,
40    /// Soft-capping of attention logits / final logits (gemma-2 family).
41    pub attn_logit_softcap: Option<f64>,
42    pub final_logit_softcap: Option<f64>,
43    /// Attention scale divisor override (`1/sqrt(query_pre_attn_scalar)`
44    /// instead of `1/sqrt(head_dim)`).
45    pub query_pre_attn_scalar: Option<f64>,
46    /// Global-layer RoPE base + scaling.
47    pub rope_theta: f64,
48    pub rope_scaling: RopeScaling,
49    /// Sliding layers rotate with this base instead (gemma dual-RoPE).
50    pub rope_local_theta: Option<f64>,
51    /// One entry per transformer layer.
52    pub layers: Vec<LayerKind>,
53}
54
55impl ArchSpec {
56    pub fn resolve(meta: &ModelMetadata) -> Self {
57        let pattern = &meta.attention_pattern;
58        let n = meta.num_hidden_layers;
59        let arch = meta.architecture.as_str();
60
61        let layers: Vec<LayerKind> = match (arch, pattern.sliding_window) {
62            (_, None) => vec![LayerKind::Global; n],
63            // Gemma interleave: every `pattern`-th layer is global.
64            ("gemma3" | "gemma3_text", Some(w)) => (0..n)
65                .map(|i| {
66                    if pattern.is_global_layer(i) {
67                        LayerKind::Global
68                    } else {
69                        LayerKind::Sliding(w)
70                    }
71                })
72                .collect(),
73            // Qwen2/3 partition: the first `max_window_layers` layers are
74            // full attention, the rest slide. HF defaults the split to the
75            // layer count, i.e. nothing slides.
76            ("qwen2" | "qwen3", Some(w)) => {
77                let full = pattern.max_window_layers.unwrap_or(n);
78                (0..n)
79                    .map(|i| if i < full { LayerKind::Global } else { LayerKind::Sliding(w) })
80                    .collect()
81            }
82            // Mistral v0.1 and phi-3: every layer slides (phi-3-mini ships
83            // sliding_window 2047 alongside its 4k context).
84            ("mistral" | "phi3", Some(w)) => vec![LayerKind::Sliding(w); n],
85            // Unknown sliding semantics: run global rather than guess a
86            // wrong mask (the registry guards refuse the risky cases).
87            (_, Some(_)) => vec![LayerKind::Global; n],
88        };
89
90        let gemma = matches!(arch, "gemma3" | "gemma3_text");
91        ArchSpec {
92            // Gemma is forced: its GGUFs carry no activation key and the
93            // family never uses silu.
94            activation: if gemma { Activation::GeluTanh } else { meta.activation },
95            norm_flavor: if gemma { NormFlavor::GemmaRmsNorm } else { NormFlavor::RmsNorm },
96            qk_norm: matches!(arch, "gemma3" | "gemma3_text" | "qwen3"),
97            sandwich_norms: gemma,
98            embed_scale_sqrt_hidden: gemma,
99            attn_logit_softcap: None,
100            final_logit_softcap: None,
101            query_pre_attn_scalar: pattern.query_pre_attn_scalar,
102            rope_theta: meta.rope_theta,
103            rope_scaling: meta.rope_scaling.clone(),
104            rope_local_theta: (gemma && pattern.sliding_window.is_some())
105                .then_some(pattern.rope_local_theta),
106            layers,
107        }
108    }
109
110    /// Per-layer window sizes in the shape `PagedKVCache::new_with_windows`
111    /// takes (`None` = global arena, `Some(w)` = rolling window).
112    pub fn windows(&self) -> Vec<Option<usize>> {
113        self.layers
114            .iter()
115            .map(|k| match k {
116                LayerKind::Global => None,
117                LayerKind::Sliding(w) => Some(*w),
118            })
119            .collect()
120    }
121}
122
123#[cfg(test)]
124mod tests {
125    use super::*;
126    use combs_formats::AttentionPattern;
127
128    fn meta(arch: &str, layers: usize, pattern: AttentionPattern) -> ModelMetadata {
129        let mut m = ModelMetadata::diffusion_placeholder(arch);
130        m.num_hidden_layers = layers;
131        m.attention_pattern = pattern;
132        m
133    }
134
135    #[test]
136    fn gemma_every_nth_layer_is_global() {
137        let p = AttentionPattern {
138            sliding_window: Some(512),
139            pattern: 6,
140            ..AttentionPattern::default()
141        };
142        let spec = ArchSpec::resolve(&meta("gemma3_text", 12, p));
143        for (i, k) in spec.layers.iter().enumerate() {
144            let expect = if (i + 1) % 6 == 0 { LayerKind::Global } else { LayerKind::Sliding(512) };
145            assert_eq!(*k, expect, "layer {i}");
146        }
147        assert!(spec.qk_norm && spec.sandwich_norms && spec.embed_scale_sqrt_hidden);
148        assert_eq!(spec.norm_flavor, NormFlavor::GemmaRmsNorm);
149        assert_eq!(spec.rope_local_theta, Some(10_000.0));
150    }
151
152    #[test]
153    fn qwen_first_n_layers_are_global() {
154        let p = AttentionPattern {
155            sliding_window: Some(4096),
156            max_window_layers: Some(2),
157            ..AttentionPattern::default()
158        };
159        let spec = ArchSpec::resolve(&meta("qwen2", 4, p));
160        assert_eq!(
161            spec.layers,
162            vec![
163                LayerKind::Global,
164                LayerKind::Global,
165                LayerKind::Sliding(4096),
166                LayerKind::Sliding(4096)
167            ]
168        );
169        // Default split = layer count: nothing slides.
170        let p = AttentionPattern {
171            sliding_window: Some(4096),
172            ..AttentionPattern::default()
173        };
174        let spec = ArchSpec::resolve(&meta("qwen2", 4, p));
175        assert_eq!(spec.layers, vec![LayerKind::Global; 4]);
176    }
177
178    #[test]
179    fn mistral_slides_everywhere_and_llama_stays_global() {
180        let p = AttentionPattern {
181            sliding_window: Some(4096),
182            ..AttentionPattern::default()
183        };
184        let spec = ArchSpec::resolve(&meta("mistral", 3, p));
185        assert_eq!(spec.layers, vec![LayerKind::Sliding(4096); 3]);
186        assert_eq!(spec.windows(), vec![Some(4096); 3]);
187
188        let spec = ArchSpec::resolve(&meta("llama", 3, AttentionPattern::default()));
189        assert_eq!(spec.layers, vec![LayerKind::Global; 3]);
190        assert!(!spec.qk_norm && !spec.sandwich_norms);
191        assert_eq!(spec.norm_flavor, NormFlavor::RmsNorm);
192    }
193}